Enh/add time agregation selection #204
7 changed files with 295 additions and 4 deletions
|
|
@ -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
|
||||
~~~~
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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").
|
||||
|
or,
`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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
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.
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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in a new issue
method=>or,