diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index a63d2bfa6..9eef40a05 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -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 ~~~~ diff --git a/junifer/markers/parcel_aggregation.py b/junifer/markers/parcel_aggregation.py index 9cb5b69d6..c8c3f892c 100644 --- a/junifer/markers/parcel_aggregation.py +++ b/junifer/markers/parcel_aggregation.py @@ -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 ` 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 diff --git a/junifer/markers/sphere_aggregation.py b/junifer/markers/sphere_aggregation.py index 756e1d189..1a30bccdb 100644 --- a/junifer/markers/sphere_aggregation.py +++ b/junifer/markers/sphere_aggregation.py @@ -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"). 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 ` 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 diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index e54bba8ef..84d5d4c58 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -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) diff --git a/junifer/markers/tests/test_sphere_aggregation.py b/junifer/markers/tests/test_sphere_aggregation.py index 285635058..68ea96714 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -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) diff --git a/junifer/stats.py b/junifer/stats.py index 8e0ef02cf..06862335e 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -4,7 +4,7 @@ # Synchon Mandal # 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) 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 diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py index 8567b72bd..c8b5fce63 100644 --- a/junifer/tests/test_stats.py +++ b/junifer/tests/test_stats.py @@ -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)