From fa8fcfe9357ecd574f0683e75d2f6a838351d878 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 26 Feb 2025 22:36:19 +0100 Subject: [PATCH 1/7] Add TemporalFilter --- .../pipeline/pipeline_component_registry.py | 1 + junifer/preprocess/__init__.pyi | 2 + junifer/preprocess/filter/__init__.py | 9 + junifer/preprocess/filter/__init__.pyi | 3 + junifer/preprocess/filter/temporalfilter.py | 243 ++++++++++++++++++ 5 files changed, 258 insertions(+) create mode 100644 junifer/preprocess/filter/__init__.py create mode 100644 junifer/preprocess/filter/__init__.pyi create mode 100644 junifer/preprocess/filter/temporalfilter.py diff --git a/junifer/pipeline/pipeline_component_registry.py b/junifer/pipeline/pipeline_component_registry.py index 55a8fa31a..7ac530b37 100644 --- a/junifer/pipeline/pipeline_component_registry.py +++ b/junifer/pipeline/pipeline_component_registry.py @@ -76,6 +76,7 @@ class PipelineComponentRegistry(metaclass=Singleton): "SpaceWarper": "SpaceWarper", "fMRIPrepConfoundRemover": "fMRIPrepConfoundRemover", "TemporalSlicer": "TemporalSlicer", + "TemporalFilter": "TemporalFilter", }, "marker": { "ALFFParcels": "ALFFParcels", diff --git a/junifer/preprocess/__init__.pyi b/junifer/preprocess/__init__.pyi index c62792917..4cff4dd78 100644 --- a/junifer/preprocess/__init__.pyi +++ b/junifer/preprocess/__init__.pyi @@ -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 .filter import TemporalFilter diff --git a/junifer/preprocess/filter/__init__.py b/junifer/preprocess/filter/__init__.py new file mode 100644 index 000000000..d63e82a5b --- /dev/null +++ b/junifer/preprocess/filter/__init__.py @@ -0,0 +1,9 @@ +"""Provide imports for filter sub-package.""" + +# Authors: Federico Raimondo +# License: AGPL + +import lazy_loader as lazy + + +__getattr__, __dir__, __all__ = lazy.attach_stub(__name__, __file__) diff --git a/junifer/preprocess/filter/__init__.pyi b/junifer/preprocess/filter/__init__.pyi new file mode 100644 index 000000000..685aa3091 --- /dev/null +++ b/junifer/preprocess/filter/__init__.pyi @@ -0,0 +1,3 @@ +__all__ = ["TemporalFilter"] + +from .temporalfilter import TemporalFilter diff --git a/junifer/preprocess/filter/temporalfilter.py b/junifer/preprocess/filter/temporalfilter.py new file mode 100644 index 000000000..a03d05768 --- /dev/null +++ b/junifer/preprocess/filter/temporalfilter.py @@ -0,0 +1,243 @@ +"""Provide class for temporal filtering.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# 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 ` 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}") + + # Deconfound data + 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 deconfounded 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 -- 2.52.0 From 4168e6606dbb3eed7b25dff4f8e3a348c9269aed Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 21 Jul 2025 13:45:59 +0200 Subject: [PATCH 2/7] refactor: move TemporalFilter to reduce file layout complexity --- junifer/preprocess/__init__.pyi | 2 +- .../temporalfilter.py => _temporal_filter.py} | 14 ++++++-------- junifer/preprocess/filter/__init__.py | 9 --------- junifer/preprocess/filter/__init__.pyi | 3 --- 4 files changed, 7 insertions(+), 21 deletions(-) rename junifer/preprocess/{filter/temporalfilter.py => _temporal_filter.py} (96%) delete mode 100644 junifer/preprocess/filter/__init__.py delete mode 100644 junifer/preprocess/filter/__init__.pyi diff --git a/junifer/preprocess/__init__.pyi b/junifer/preprocess/__init__.pyi index 4cff4dd78..29649c62d 100644 --- a/junifer/preprocess/__init__.pyi +++ b/junifer/preprocess/__init__.pyi @@ -12,4 +12,4 @@ from .confounds import fMRIPrepConfoundRemover from .warping import SpaceWarper from .smoothing import Smoothing from ._temporal_slicer import TemporalSlicer -from .filter import TemporalFilter +from ._temporal_filter import TemporalFilter diff --git a/junifer/preprocess/filter/temporalfilter.py b/junifer/preprocess/_temporal_filter.py similarity index 96% rename from junifer/preprocess/filter/temporalfilter.py rename to junifer/preprocess/_temporal_filter.py index a03d05768..d441052ae 100644 --- a/junifer/preprocess/filter/temporalfilter.py +++ b/junifer/preprocess/_temporal_filter.py @@ -16,12 +16,12 @@ 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 +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"] @@ -126,8 +126,6 @@ class TemporalFilter(BasePreprocessor): # BOLD must be 4D niimg check_niimg_4d(input["data"]) - - def preprocess( self, input: dict[str, Any], diff --git a/junifer/preprocess/filter/__init__.py b/junifer/preprocess/filter/__init__.py deleted file mode 100644 index d63e82a5b..000000000 --- a/junifer/preprocess/filter/__init__.py +++ /dev/null @@ -1,9 +0,0 @@ -"""Provide imports for filter sub-package.""" - -# Authors: Federico Raimondo -# License: AGPL - -import lazy_loader as lazy - - -__getattr__, __dir__, __all__ = lazy.attach_stub(__name__, __file__) diff --git a/junifer/preprocess/filter/__init__.pyi b/junifer/preprocess/filter/__init__.pyi deleted file mode 100644 index 685aa3091..000000000 --- a/junifer/preprocess/filter/__init__.pyi +++ /dev/null @@ -1,3 +0,0 @@ -__all__ = ["TemporalFilter"] - -from .temporalfilter import TemporalFilter -- 2.52.0 From 218b4e9d01de3fd8548979078d00852f157d6788 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 21 Jul 2025 13:49:52 +0200 Subject: [PATCH 3/7] update: add tests for TemporalFilter --- .../preprocess/tests/test_temporal_filter.py | 92 +++++++++++++++++++ 1 file changed, 92 insertions(+) create mode 100644 junifer/preprocess/tests/test_temporal_filter.py diff --git a/junifer/preprocess/tests/test_temporal_filter.py b/junifer/preprocess/tests/test_temporal_filter.py new file mode 100644 index 000000000..c6d63b999 --- /dev/null +++ b/junifer/preprocess/tests/test_temporal_filter.py @@ -0,0 +1,92 @@ +"""Provide tests for TemporalFilter.""" + +# Authors: Synchon Mandal +# 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", + ( + [ + True, + True, + None, + None, + None, + ], + [ + False, + True, + 0.1, + None, + None, + ], + [ + True, + False, + None, + 0.08, + None, + ], + [ + False, + False, + None, + None, + 2, + ], + [ + True, + True, + 0.1, + 0.08, + 2, + ], + ), +) +def test_TemporalFilter( + detrend: bool, + standardize: bool, + low_pass: Optional[float], + high_pass: Optional[float], + t_r: Optional[float], +) -> 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. + expect : typing.ContextManager + The parametrized ContextManager object. + + """ + 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, + ).fit_transform(element_data) + + assert isinstance(output, dict) -- 2.52.0 From af8b1720ab9474a6817949b02a5720e0a77cc971 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 21 Jul 2025 13:55:47 +0200 Subject: [PATCH 4/7] chore: add changelog 432.feature --- docs/changes/newsfragments/432.feature | 1 + 1 file changed, 1 insertion(+) create mode 100644 docs/changes/newsfragments/432.feature diff --git a/docs/changes/newsfragments/432.feature b/docs/changes/newsfragments/432.feature new file mode 100644 index 000000000..c23702d3d --- /dev/null +++ b/docs/changes/newsfragments/432.feature @@ -0,0 +1 @@ +Introduce :class:`.TemporalFilter` preprocessor for temporally filtering BOLD data by `Fede Raimondo`_ -- 2.52.0 From c714d2347fa3b4b0e60c44b788ee6fb7390b09a9 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 21 Jul 2025 15:05:31 +0200 Subject: [PATCH 5/7] update: improve tests for TemporalFilter --- junifer/preprocess/tests/test_temporal_filter.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/junifer/preprocess/tests/test_temporal_filter.py b/junifer/preprocess/tests/test_temporal_filter.py index c6d63b999..7efb64a7d 100644 --- a/junifer/preprocess/tests/test_temporal_filter.py +++ b/junifer/preprocess/tests/test_temporal_filter.py @@ -13,7 +13,7 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber @pytest.mark.parametrize( - "detrend, standardize, low_pass, high_pass, t_r", + "detrend, standardize, low_pass, high_pass, t_r, masks", ( [ True, @@ -21,6 +21,7 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber None, None, None, + None, ], [ False, @@ -28,6 +29,7 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber 0.1, None, None, + "compute_brain_mask", ], [ True, @@ -35,6 +37,7 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber None, 0.08, None, + "compute_background_mask", ], [ False, @@ -42,6 +45,7 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber None, None, 2, + None, ], [ True, @@ -49,6 +53,7 @@ from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber 0.1, 0.08, 2, + "compute_brain_mask", ], ), ) @@ -58,6 +63,7 @@ def test_TemporalFilter( low_pass: Optional[float], high_pass: Optional[float], t_r: Optional[float], + masks: Optional[str], ) -> None: """Test TemporalFilter. @@ -73,8 +79,8 @@ def test_TemporalFilter( The parametrized high pass value. t_r : float or None The parametrized repetition time. - expect : typing.ContextManager - The parametrized ContextManager object. + masks : str or None + The parametrized mask. """ with PartlyCloudyTestingDataGrabber() as dg: @@ -87,6 +93,7 @@ def test_TemporalFilter( low_pass=low_pass, high_pass=high_pass, t_r=t_r, + masks=masks, ).fit_transform(element_data) assert isinstance(output, dict) -- 2.52.0 From 18b919877d649d81aa4f24766a06787e0d9d4872 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 21 Jul 2025 15:07:39 +0200 Subject: [PATCH 6/7] docs: add TemporalFilter to builtin.rst --- docs/builtin.rst | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/docs/builtin.rst b/docs/builtin.rst index e802ab47b..e67fd7a32 100644 --- a/docs/builtin.rst +++ b/docs/builtin.rst @@ -167,6 +167,10 @@ Available - Slice ``BOLD`` data temporally - | Done - :gh:`443` + * - ``TemporalFilter`` + - Filter (clean) ``BOLD`` data temporally + - | Done + - :gh:`432` .. -- 2.52.0 From 311602dacaa3a022b1953d7e7c00780ba831b631 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 21 Jul 2025 15:10:27 +0200 Subject: [PATCH 7/7] chore: improve commentary in TemporalFilter --- junifer/preprocess/_temporal_filter.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/junifer/preprocess/_temporal_filter.py b/junifer/preprocess/_temporal_filter.py index d441052ae..517bfd6ca 100644 --- a/junifer/preprocess/_temporal_filter.py +++ b/junifer/preprocess/_temporal_filter.py @@ -211,7 +211,6 @@ class TemporalFilter(BasePreprocessor): logger.debug(f"\thigh_pass: {self.high_pass}") logger.debug(f"\tt_r: {self.t_r}") - # Deconfound data cleaned_img = nimg.clean_img( imgs=bold_img, detrend=self.detrend, @@ -224,7 +223,7 @@ class TemporalFilter(BasePreprocessor): ) # Fix t_r as nilearn messes it up cleaned_img.header["pixdim"][4] = t_r - # Save deconfounded data + # Save filtered data filtered_data_path = element_tempdir / "filtered_data.nii.gz" nib.save(cleaned_img, filtered_data_path) -- 2.52.0