[ENH]: Add TemporalSlicer #443
6 changed files with 382 additions and 0 deletions
|
|
@ -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`
|
||||||
|
|
||||||
|
|
||||||
..
|
..
|
||||||
|
|
|
||||||
1
docs/changes/newsfragments/443.feature
Normal file
1
docs/changes/newsfragments/443.feature
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Introduce :class:`.TemporalSlicer` preprocessor for temporally slicing BOLD data by `Synchon Mandal`_
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
236
junifer/preprocess/_temporal_slicer.py
Normal file
236
junifer/preprocess/_temporal_slicer.py
Normal 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 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
|
||||||
138
junifer/preprocess/tests/test_temporal_slicer.py
Normal file
138
junifer/preprocess/tests/test_temporal_slicer.py
Normal 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
|
||||||
Loading…
Reference in a new issue
Can we thorugh out of bounds exceptions ourselves?
Maybe also an INFO log to show the user what's being "sliced".