[ENH]: Add TemporalSlicer #443

Merged
synchon merged 17 commits from feat/temporal-slicer-preproc into main 2025-05-09 08:30:25 +00:00
6 changed files with 382 additions and 0 deletions

View file

@ -163,6 +163,10 @@ Available
| ``fMRIPrep``-ed data | ``fMRIPrep``-ed data
- In Progress - In Progress
- :gh:`161` - :gh:`161`
* - ``TemporalSlicer``
- Slice ``BOLD`` data temporally
- | Done
- :gh:`443`
.. ..

View file

@ -0,0 +1 @@
Introduce :class:`.TemporalSlicer` preprocessor for temporally slicing BOLD data by `Synchon Mandal`_

View file

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

View file

@ -3,9 +3,11 @@ __all__ = [
"fMRIPrepConfoundRemover", "fMRIPrepConfoundRemover",
"SpaceWarper", "SpaceWarper",
"Smoothing", "Smoothing",
"TemporalSlicer",
] ]
from .base import BasePreprocessor from .base import BasePreprocessor
from .confounds import fMRIPrepConfoundRemover from .confounds import fMRIPrepConfoundRemover
from .warping import SpaceWarper from .warping import SpaceWarper
from .smoothing import Smoothing from .smoothing import Smoothing
from ._temporal_slicer import TemporalSlicer

View file

@ -0,0 +1,236 @@
"""Provide class for temporal slicing."""
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Any, ClassVar, Optional
import nibabel as nib
import nilearn.image as nimg
from ..api.decorators import register_preprocessor
from ..pipeline import WorkDirManager
from ..typing import Dependencies
from ..utils import logger, raise_error
from .base import BasePreprocessor
__all__ = ["TemporalSlicer"]
@register_preprocessor
class TemporalSlicer(BasePreprocessor):
"""Class for temporal slicing.
Parameters
----------
start : float
Starting time point, in second.
stop : float or None
Ending time point, in second. If None, stops at the last time point.
Can also do negative indexing and has the same meaning as standard
Python slicing except it represents time points.
duration : float or None, optional
Time duration to add to ``start``, in second. If None, ``stop`` is
respected, else error is raised (default None).
t_r : float or None, optional
Repetition time, in second (sampling period).
If None, it will use t_r from nifti header (default None).
Raises
------
ValueError
If ``start`` is negative.
"""
_DEPENDENCIES: ClassVar[Dependencies] = {"nilearn"}
def __init__(
self,
start: float,
stop: Optional[float],
duration: Optional[float] = None,
t_r: Optional[float] = None,
) -> None:
"""Initialize the class."""
if start < 0:
raise_error("`start` cannot be negative")
else:
self.start = start
self.stop = stop
self.duration = duration
self.t_r = t_r
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 data type input to the preprocessor.
Returns
-------
str
The data type output by the preprocessor.
"""
# Does not add any new keys
return input_type
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
The input from the Junifer Data object.
extra_input : dict, optional
The other fields in the Junifer Data object.
Returns
-------
dict
The computed result as dictionary.
None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
Raises
------
RuntimeError
If no time slicing will be performed or
if ``stop`` is not None when ``duration`` is provided or
if calculated stop index is greater than allowed value.
"""
logger.debug("Temporal slicing")
# Get BOLD data
bold_img = input["data"]
time_dim = bold_img.shape[3]
# Check if slicing is not required
if self.start == 0:
if self.stop is None or self.stop == -1 or self.stop == time_dim:
raise_error(
"No temporal slicing will be performed as "
f"`start` = {self.start} and "
f"`stop` = {self.stop}, hence you "
"should remove the TemporalSlicer from the preprocess "
"step.",
klass=RuntimeError,
)
# Sanity check for stop and duration combination
if self.duration is not None and self.stop is not None:
raise_error(
"`stop` should be None if `duration` is not None. "
"Set `stop` = None for TemporalSlicer to continue.",
klass=RuntimeError,
)
# 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
element_tempdir = WorkDirManager().get_element_tempdir(
prefix="temporal_slicer"
)
# Check stop; duration is None
if self.stop is None:
if self.duration is not None:
stop = self.start + self.duration
else:
stop = time_dim
else:
# Calculate stop index if going from end
if self.stop < 0:
stop = time_dim + 1 + self.stop
else:
stop = self.stop
# Convert slice range from seconds to indices
index = slice(int(self.start // t_r), int(stop // t_r))
fraimondo commented 2025-04-11 11:29:38 +00:00 (Migrated from github.com)

Can we thorugh out of bounds exceptions ourselves?

Maybe also an INFO log to show the user what's being "sliced".

Can we thorugh out of bounds exceptions ourselves? Maybe also an INFO log to show the user what's being "sliced".
synchon commented 2025-04-14 09:26:48 +00:00 (Migrated from github.com)

@fraimondo The new commits should address your comments, feel free to add an approval if you are ok with the PR.

@fraimondo The new commits should address your comments, feel free to add an approval if you are ok with the PR.
# Check if stop index is out of bounds
if index.stop > time_dim:
raise_error(
f"Calculated stop index: {index.stop} is greater than "
f"allowed value: {time_dim}",
klass=IndexError,
)
logger.info(
"Computed slice range for TemporalSlicer: "
f"[{index.start},{index.stop}]"
)
# Slice image
sliced_img = nimg.index_img(bold_img, index)
# Fix t_r as nilearn messes it up
sliced_img.header["pixdim"][4] = t_r
# Save sliced data
sliced_img_path = element_tempdir / "sliced_data.nii.gz"
nib.save(sliced_img, sliced_img_path)
logger.debug("Updating `BOLD`")
input.update(
{
# Update path to sync with "data"
"path": sliced_img_path,
# Update data
"data": sliced_img,
}
)
# Check for BOLD.confounds and update if found
if input.get("confounds") is not None:
# Slice confounds
sliced_confounds_df = input["confounds"]["data"].iloc[index, :]
# Save sliced confounds
sliced_confounds_path = (
element_tempdir / "sliced_confounds_regressors.tsv"
)
sliced_confounds_df.to_csv(
sliced_confounds_path,
sep="\t",
index=False,
)
logger.debug("Updating `BOLD.confounds`")
input["confounds"].update(
{
# Update path to sync with "data"
"path": sliced_confounds_path,
# Update data
"data": sliced_confounds_df,
}
)
return input, None

View file

@ -0,0 +1,138 @@
"""Provide tests for TemporalSlicer."""
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from contextlib import AbstractContextManager, nullcontext
from typing import Optional
import pytest
from junifer.datareader import DefaultDataReader
from junifer.preprocess import TemporalSlicer
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
@pytest.mark.parametrize(
"start, stop, duration, t_r, expected_dim, expect",
(
[
0.0,
168.0,
None,
2.0,
84,
pytest.raises(RuntimeError, match="No temporal slicing"),
], # t_r from doc is 2.0
[0.0, 167.0, None, 2.0, 83, nullcontext()], # t_r from doc is 2.0
[
0.0,
None,
None,
2.0,
84,
pytest.raises(RuntimeError, match="No temporal slicing"),
], # total with no end
[2.0, None, None, 2.0, 83, nullcontext()], # total with no end
[0.0, 84.0, None, 2.0, 42, nullcontext()], # first half
[0.0, -85.0, None, 2.0, 42, nullcontext()], # first half from end
[84.0, -1.0, None, 2.0, 42, nullcontext()], # second half
[
33.0,
-33.0,
33.0,
2.0,
42,
pytest.raises(RuntimeError, match="`stop` should be None"),
],
[10.0, None, 30.0, 2.0, 15, nullcontext()],
[
0.0,
168.0,
None,
None,
168,
pytest.raises(RuntimeError, match="No temporal slicing"),
], # t_r from image is 1.0
[0.0, 167.0, None, None, 167, nullcontext()], # t_r from image is 1.0
[
0.0,
None,
None,
None,
168,
pytest.raises(RuntimeError, match="No temporal slicing"),
], # total with no end
[0.0, 84.0, None, None, 84, nullcontext()], # first half
[0.0, -85.0, None, None, 84, nullcontext()], # first half from end
[84.0, -1.0, None, None, 84, nullcontext()], # second half
[
33.0,
-33.0,
33.0,
None,
84,
pytest.raises(RuntimeError, match="`stop` should be None"),
],
[10.0, None, 30.0, None, 30, nullcontext()],
[
-1.0,
None,
None,
None,
84,
pytest.raises(ValueError, match="`start` cannot be negative"),
],
[
0.0,
500.0,
None,
2.0,
42,
pytest.raises(IndexError, match="Calculated stop index:"),
],
),
)
def test_TemporalSlicer(
start: float,
stop: Optional[float],
duration: Optional[float],
t_r: Optional[float],
expected_dim: int,
expect: AbstractContextManager,
) -> None:
"""Test TemporalSlicer.
Parameters
----------
start : float
The parametrized start.
stop : float or None
The parametrized stop.
duration : float or None
The parametrized duration.
t_r : float or None
The parametrized TR.
expected_dim : int
The parametrized expected time dimension size.
expect : typing.ContextManager
The parametrized ContextManager object.
"""
with PartlyCloudyTestingDataGrabber() as dg:
# Read data
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Preprocess data
with expect:
output = TemporalSlicer(
start=start, # in seconds
stop=stop, # in seconds
duration=duration, # in seconds
t_r=t_r, # in seconds
).fit_transform(element_data)
# Check image data dim
assert output["BOLD"]["data"].shape[3] == expected_dim
# Check confounds dim
assert output["BOLD"]["confounds"]["data"].shape[0] == expected_dim