Enh/add time agregation selection #204

Merged
fraimondo merged 7 commits from enh/add_time_agregation_selection into main 2023-03-28 09:41:13 +00:00
7 changed files with 295 additions and 4 deletions

View file

@ -71,6 +71,10 @@ Enhancements
- Add copy button to documentation code blocks (:gh:`205` by `Synchon Mandal`_).
- Add func:`junifer.stats.select` as an aggregation function that allows to select a subset of elements (:gh:`204` by `Fede Raimondo`_).
- Add ``time_method`` and ``time_method_params`` to :class:`junifer.markers.ParcelAggregation` and :class:`junifer.markers.SphereAggregation`, allowing to apply an aggregation on the time axis after the aggregation on the parcels and spheres respectively (:gh:`204` by `Fede Raimondo`_).
Bugs
~~~~

View file

@ -13,7 +13,7 @@ from nilearn.maskers import NiftiMasker
from ..api.decorators import register_marker
from ..data import get_mask, load_parcellation, merge_parcellations
from ..stats import get_aggfunc_by_name
from ..utils import logger
from ..utils import logger, warn_with_log, raise_error
from .base import BaseMarker
@ -32,6 +32,12 @@ class ParcelAggregation(BaseMarker):
method_params : dict, optional
synchon commented 2023-03-27 15:27:06 +00:00 (Migrated from github.com)

method =>

``method``

or,

:term:`method`
`method` => ``` ``method`` ``` or, ``` :term:`method` ```
Parameters to pass to the aggregation function. Check valid options in
:func:`junifer.stats.get_aggfunc_by_name`.
time_method : str, optional
The method to use to aggregate the time series over the time points,
after applying :term:`method` (only applicable to BOLD data). If None,
it will not operate on the time dimension (default None).
time_method_params : dict, optional
The parameters to pass to the time aggregation method (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.
@ -52,6 +58,8 @@ class ParcelAggregation(BaseMarker):
parcellation: Union[str, List[str]],
method: str,
method_params: Optional[Dict[str, Any]] = None,
time_method: Optional[str] = None,
time_method_params: Optional[Dict[str, Any]] = None,
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
on: Union[List[str], str, None] = None,
name: Optional[str] = None,
@ -64,6 +72,20 @@ class ParcelAggregation(BaseMarker):
self.masks = masks
super().__init__(on=on, name=name)
# Verify after super init so self._on is set
if "BOLD" not in self._on and time_method is not None:
raise_error(
"`time_method` can only be used with BOLD data. "
"Please remove `time_method` parameter."
)
if time_method is None and time_method_params is not None:
raise_error(
"`time_method_params` can only be used with `time_method`. "
"Please remove `time_method_params` parameter."
)
self.time_method = time_method
self.time_method_params = time_method_params or {}
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
@ -191,5 +213,18 @@ class ParcelAggregation(BaseMarker):
# in it
out_values = np.array(out_values).T
if self.time_method is not None:
if out_values.shape[0] > 1:
logger.debug("Aggregating time dimension")
time_agg_func = get_aggfunc_by_name(
self.time_method, func_params=self.time_method_params
)
out_values = time_agg_func(out_values, axis=0)
else:
warn_with_log(
"No time dimension to aggregate as only one time point is "
"available."
)
out = {"data": out_values, "col_names": labels}
return out

View file

@ -10,7 +10,7 @@ from ..api.decorators import register_marker
from ..data import get_mask, load_coordinates
from ..external.nilearn import JuniferNiftiSpheresMasker
from ..stats import get_aggfunc_by_name
from ..utils import logger
from ..utils import logger, raise_error, warn_with_log
from .base import BaseMarker
@ -37,6 +37,12 @@ class SphereAggregation(BaseMarker):
(default "mean").
synchon commented 2023-03-27 15:34:39 +00:00 (Migrated from github.com)

method =>

``method``

or,

:term:`method`
`method` => ``` ``method`` ``` or, ``` :term:`method` ```
method_params : dict, optional
The parameters to pass to the aggregation method (default None).
time_method : str, optional
The method to use to aggregate the time series over the time points,
after applying :term:`method` (only applicable to BOLD data). If None,
it will not operate on the time dimension (default None).
time_method_params : dict, optional
The parameters to pass to the time aggregation method (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.
@ -60,6 +66,8 @@ class SphereAggregation(BaseMarker):
allow_overlap: bool = False,
method: str = "mean",
method_params: Optional[Dict[str, Any]] = None,
time_method: Optional[str] = None,
time_method_params: Optional[Dict[str, Any]] = None,
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
on: Union[List[str], str, None] = None,
name: Optional[str] = None,
@ -72,6 +80,20 @@ class SphereAggregation(BaseMarker):
self.masks = masks
super().__init__(on=on, name=name)
# Verify after super init so self._on is set
if "BOLD" not in self._on and time_method is not None:
raise_error(
"`time_method` can only be used with BOLD data. "
"Please remove `time_method` parameter."
)
if time_method is None and time_method_params is not None:
raise_error(
"`time_method_params` can only be used with `time_method`. "
"Please remove `time_method_params` parameter."
)
self.time_method = time_method
self.time_method_params = time_method_params or {}
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
@ -158,6 +180,18 @@ class SphereAggregation(BaseMarker):
)
# Fit and transform the marker on the data
out_values = masker.fit_transform(t_input_img)
if self.time_method is not None:
if out_values.shape[0] > 1:
logger.debug("Aggregating time dimension")
time_agg_func = get_aggfunc_by_name(
self.time_method, func_params=self.time_method_params
)
out_values = time_agg_func(out_values, axis=0)
else:
warn_with_log(
"No time dimension to aggregate as only one time point is "
"available."
)
# Format the output
out = {"data": out_values, "col_names": out_labels}
return out

View file

@ -556,3 +556,69 @@ def test_ParcelAggregation_3D_multiple_duplicated_labels(
col_names = [f"Schaefer100x7_low_{x}" for x in labels1]
col_names += [f"Schaefer100x7_high_{x}" for x in labels2]
assert col_names == split_mean["col_names"]
def test_ParcelAggregation_4D_agg_time():
"""Test ParcelAggregation object on 4D images, aggregating time."""
# Get the testing parcellation (for nilearn)
parcellation = datasets.fetch_atlas_schaefer_2018(
n_rois=100, yeo_networks=7, resolution_mm=2
)
# Get the SPM auditory data:
subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore
# Create NiftiLabelsMasker
nifti_masker = NiftiLabelsMasker(labels_img=parcellation.maps)
auto4d = nifti_masker.fit_transform(fmri_img)
auto_mean = auto4d.mean(axis=0)
# Create ParcelAggregation object
marker = ParcelAggregation(
parcellation="Schaefer100x7", method="mean", time_method="mean"
)
input = {"BOLD": {"data": fmri_img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 1
assert_array_equal(auto_mean.shape, jun_values4d.shape)
assert_array_almost_equal(auto_mean, jun_values4d, decimal=2)
auto_pick_0 = auto4d[:1, :]
marker = ParcelAggregation(
parcellation="Schaefer100x7",
method="mean",
time_method="select",
time_method_params={"pick": [0]},
)
input = {"BOLD": {"data": fmri_img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 2
assert_array_equal(auto_pick_0.shape, jun_values4d.shape)
assert_array_equal(auto_pick_0, jun_values4d)
with pytest.raises(ValueError, match="can only be used with BOLD data"):
ParcelAggregation(
parcellation="Schaefer100x7",
method="mean",
time_method="select",
time_method_params={"pick": [0]},
on="VBM_GM",
)
with pytest.raises(
ValueError, match="can only be used with `time_method`"
):
ParcelAggregation(
parcellation="Schaefer100x7",
method="mean",
time_method_params={"pick": [0]},
on="VBM_GM",
)
with pytest.warns(RuntimeWarning, match="No time dimension to aggregate"):
input = {"BOLD": {"data": fmri_img.slicer[..., 0:1], "meta": {}}}
marker.fit_transform(input)

View file

@ -167,3 +167,70 @@ def test_SphereAggregation_3D_mask() -> None:
assert jun_values4d.ndim == 2
assert_array_equal(auto4d.shape, jun_values4d.shape)
assert_array_equal(auto4d, jun_values4d)
def test_SphereAggregation_4D_agg_time() -> None:
"""Test SphereAggregation object on 4D images, aggregating time."""
# Get the testing coordinates (for nilearn)
coordinates, _ = load_coordinates(COORDS)
# Get the SPM auditory data
subject_data = datasets.fetch_spm_auditory()
fmri_img = concat_imgs(subject_data.func) # type: ignore
# Create NiftSpheresMasker
nifti_masker = NiftiSpheresMasker(seeds=coordinates, radius=RADIUS)
auto4d = nifti_masker.fit_transform(fmri_img)
auto_mean = auto4d.mean(axis=0)
# Create SphereAggregation object
marker = SphereAggregation(
coords=COORDS, method="mean", radius=RADIUS, time_method="mean"
)
input = {"BOLD": {"data": fmri_img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 1
assert_array_equal(auto_mean.shape, jun_values4d.shape)
assert_array_equal(auto_mean, jun_values4d)
auto_pick_0 = auto4d[:1, :]
marker = SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
time_method="select",
time_method_params={"pick": [0]},
)
input = {"BOLD": {"data": fmri_img, "meta": {}}}
jun_values4d = marker.fit_transform(input)["BOLD"]["data"]
assert jun_values4d.ndim == 2
assert_array_equal(auto_pick_0.shape, jun_values4d.shape)
assert_array_equal(auto_pick_0, jun_values4d)
with pytest.raises(ValueError, match="can only be used with BOLD data"):
SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
time_method="pick",
time_method_params={"pick": [0]},
on="VBM_GM",
)
with pytest.raises(
ValueError, match="can only be used with `time_method`"
):
SphereAggregation(
coords=COORDS,
method="mean",
radius=RADIUS,
time_method_params={"pick": [0]},
on="VBM_GM",
)
with pytest.warns(RuntimeWarning, match="No time dimension to aggregate"):
input = {"BOLD": {"data": fmri_img.slicer[..., 0:1], "meta": {}}}
marker.fit_transform(input)

View file

@ -4,7 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Any, Callable, Dict, Optional
from typing import Any, Callable, Dict, List, Optional
import numpy as np
from scipy.stats import trim_mean
@ -29,6 +29,7 @@ def get_aggfunc_by_name(
* ``std`` -> :func:`numpy.std`
* ``trim_mean`` -> :func:`scipy.stats.trim_mean`
* ``count`` -> :func:`junifer.stats.count`
* ``select`` -> :func:`junifer.stats.select`
func_params : dict, optional
Parameters to pass to the function.
@ -49,6 +50,7 @@ def get_aggfunc_by_name(
"std",
"trim_mean",
"count",
"select",
}
if func_params is None:
func_params = {}
@ -83,6 +85,14 @@ def get_aggfunc_by_name(
func = partial(trim_mean, **func_params)
elif name == "count":
func = count
elif name == "select":
pick = func_params.get("pick", None)
drop = func_params.get("drop", None)
if pick is None and drop is None:
raise_error("Either pick or drop must be specified.")
elif pick is not None and drop is not None:
raise_error("Either pick or drop must be specified, not both.")
func = partial(select, **func_params)
else:
raise_error(
f"Function {name} unknown. Please provide any of "
@ -144,3 +154,41 @@ def winsorized_mean(
win_mean = win_dat.mean(axis=axis)
synchon commented 2023-03-27 15:37:02 +00:00 (Migrated from github.com)

Might be better to explicitly mention list of what basic type we expect.

Might be better to explicitly mention list of what basic type we expect.
synchon commented 2023-03-27 15:37:08 +00:00 (Migrated from github.com)

Same as above.

Same as above.
return win_mean
def select(
data: np.ndarray,
axis: int = 0,
pick: Optional[List[int]] = None,
drop: Optional[List[int]] = None,
) -> np.ndarray:
"""Select a subset of the data.
Parameters
----------
data : numpy.ndarray
Data to select a subset from.
axis : int, optional
The axis to select a subset from (default 0).
pick : list of int, optional
List of indices to select (default None).
drop : list of int, optional
List of indices to drop (default None).
Returns
-------
numpy.ndarray
Subset of the inputted data with the select settings
applied as specified in ``select_params``.
"""
if pick is None and drop is None:
raise_error("Either pick or drop must be specified.")
elif pick is not None and drop is not None:
raise_error("Either pick or drop must be specified, not both.")
elif drop is not None:
pick = [i for i in range(data.shape[axis]) if i not in drop]
if not isinstance(pick, np.ndarray):
pick = np.array(pick) # type: ignore
out = data.take(pick, axis=axis) # type: ignore
return out

View file

@ -9,7 +9,7 @@ import numpy as np
import pytest
from numpy.testing import assert_array_equal
from junifer.stats import count, get_aggfunc_by_name, winsorized_mean
from junifer.stats import count, get_aggfunc_by_name, select, winsorized_mean
@pytest.mark.parametrize(
@ -70,6 +70,14 @@ def test_get_aggfunc_by_name_errors() -> None:
name="winsorized_mean", func_params={"limits": [0.1, 2]}
)
with pytest.raises(ValueError, match="must be specified."):
get_aggfunc_by_name(name="select", func_params=None)
with pytest.raises(ValueError, match="must be specified, not both."):
get_aggfunc_by_name(
name="select", func_params={"pick": [0], "drop": [1]}
)
def test_winsorized_mean() -> None:
"""Test winsorized mean."""
@ -94,3 +102,32 @@ def test_count() -> None:
assert_array_equal(count(input, axis=-1), ax1)
assert_array_equal(count(input, axis=1), ax1)
assert_array_equal(count(input, axis=0), ax2)
def test_select() -> None:
"""Test select."""
input = np.arange(28).reshape(7, 4)
with pytest.raises(ValueError, match="must be specified."):
select(input, axis=2)
with pytest.raises(ValueError, match="must be specified, not both."):
select(input, pick=[1], drop=[2], axis=2)
out1 = select(input, pick=[1], axis=0)
assert_array_equal(out1, input[1:2, :])
out2 = select(input, pick=[1, 4, 6], axis=0)
assert_array_equal(out2, input[[1, 4, 6], :])
out3 = select(input, drop=[0, 2, 3, 4, 5, 6], axis=0)
assert_array_equal(out1, out3)
out4 = select(input, drop=[0, 2, 3, 5], axis=0)
assert_array_equal(out2, out4)
out5 = select(input, drop=np.array([0, 2, 3, 5]), axis=0) # type: ignore
assert_array_equal(out2, out5)
out6 = select(input, pick=np.array([1, 4, 6]), axis=0) # type: ignore
assert_array_equal(out2, out6)