[ENH]: Add TemporalSlicer #443
6 changed files with 382 additions and 0 deletions
|
|
@ -163,6 +163,10 @@ Available
|
|||
| ``fMRIPrep``-ed data
|
||||
- In Progress
|
||||
- :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",
|
||||
"SpaceWarper": "SpaceWarper",
|
||||
"fMRIPrepConfoundRemover": "fMRIPrepConfoundRemover",
|
||||
"TemporalSlicer": "TemporalSlicer",
|
||||
},
|
||||
"marker": {
|
||||
"ALFFParcels": "ALFFParcels",
|
||||
|
|
|
|||
|
|
@ -3,9 +3,11 @@ __all__ = [
|
|||
"fMRIPrepConfoundRemover",
|
||||
"SpaceWarper",
|
||||
"Smoothing",
|
||||
"TemporalSlicer",
|
||||
]
|
||||
|
||||
from .base import BasePreprocessor
|
||||
from .confounds import fMRIPrepConfoundRemover
|
||||
from .warping import SpaceWarper
|
||||
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".