From 87a3bfcad145faec3911948ab1b72e1a97181dd5 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 27 Mar 2023 14:52:21 +0200 Subject: [PATCH 1/7] WIP: add select function + time aggregation option --- .../markers/tests/test_sphere_aggregation.py | 28 +++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/junifer/markers/tests/test_sphere_aggregation.py b/junifer/markers/tests/test_sphere_aggregation.py index 285635058..0fe6aa7c2 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -167,3 +167,31 @@ 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 == 2 + assert_array_equal(auto_mean.shape, jun_values4d.shape) + assert_array_equal(auto_mean, jun_values4d) + -- 2.52.0 From 51b26e741d79cf8b0dcbcda5aecddffdd55a2cfe Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 27 Mar 2023 14:53:14 +0200 Subject: [PATCH 2/7] forgot some files --- junifer/markers/sphere_aggregation.py | 30 +++++++++++++++++++++++- junifer/stats.py | 33 +++++++++++++++++++++++++++ junifer/tests/test_stats.py | 25 +++++++++++++++++++- 3 files changed, 86 insertions(+), 2 deletions(-) diff --git a/junifer/markers/sphere_aggregation.py b/junifer/markers/sphere_aggregation.py index 756e1d189..3a243fbd0 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 from .base import BaseMarker @@ -60,6 +60,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 +74,21 @@ 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 +175,17 @@ 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: + logger.debug( + "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/stats.py b/junifer/stats.py index 8e0ef02cf..66e33a2c9 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -144,3 +144,36 @@ def winsorized_mean( win_mean = win_dat.mean(axis=axis) return win_mean + + +def select(data: np.ndarray, axis: int = 0, pick=None, drop=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, optional + List of indices to select (default None). + drop : list, 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) + out = data.take(pick, axis=axis) + return out diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py index 8567b72bd..5400f412d 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( @@ -94,3 +94,26 @@ 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) -- 2.52.0 From a75e038ce94d66cacfa453652421f1e31b73c1ee Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 27 Mar 2023 16:06:38 +0200 Subject: [PATCH 3/7] Sphere aggregation with time method --- .../markers/tests/test_sphere_aggregation.py | 30 +++++++++++++++++-- junifer/stats.py | 14 ++++++++- junifer/tests/test_stats.py | 8 +++++ 3 files changed, 48 insertions(+), 4 deletions(-) diff --git a/junifer/markers/tests/test_sphere_aggregation.py b/junifer/markers/tests/test_sphere_aggregation.py index 0fe6aa7c2..661b1659d 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -185,13 +185,37 @@ def test_SphereAggregation_4D_agg_time() -> None: # Create SphereAggregation object marker = SphereAggregation( - coords=COORDS, method="mean", radius=RADIUS, - time_method="mean" + 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 == 2 + 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", + ) diff --git a/junifer/stats.py b/junifer/stats.py index 66e33a2c9..5318671d1 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -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 " @@ -146,7 +156,9 @@ def winsorized_mean( return win_mean -def select(data: np.ndarray, axis: int = 0, pick=None, drop=None) -> np.ndarray: +def select( + data: np.ndarray, axis: int = 0, pick=None, drop=None +) -> np.ndarray: """Select a subset of the data. Parameters diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py index 5400f412d..637d10984 100644 --- a/junifer/tests/test_stats.py +++ b/junifer/tests/test_stats.py @@ -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.""" -- 2.52.0 From 580a0a8206a9b91a57cc940ec22c25321adef5d7 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 27 Mar 2023 16:50:03 +0200 Subject: [PATCH 4/7] Ready to review --- docs/changes/latest.inc | 4 ++ junifer/markers/parcel_aggregation.py | 37 +++++++++++- junifer/markers/sphere_aggregation.py | 14 +++-- .../markers/tests/test_parcel_aggregation.py | 56 +++++++++++++++++++ .../markers/tests/test_sphere_aggregation.py | 4 ++ 5 files changed, 110 insertions(+), 5 deletions(-) diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index a63d2bfa6..ae7ad2827 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`_). + +- Added ``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 spheres/parcels (:gh:`204` by `Fede Raimondo`_). + Bugs ~~~~ diff --git a/junifer/markers/parcel_aggregation.py b/junifer/markers/parcel_aggregation.py index 9cb5b69d6..bfab9bd5b 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 `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 3a243fbd0..9034142f9 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, raise_error +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 `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. @@ -88,7 +94,6 @@ class SphereAggregation(BaseMarker): 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. @@ -183,9 +188,10 @@ class SphereAggregation(BaseMarker): ) out_values = time_agg_func(out_values, axis=0) else: - logger.debug( + warn_with_log( "No time dimension to aggregate as only one time point is " - "available.") + "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..cbfe138f9 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -556,3 +556,59 @@ 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.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 661b1659d..f4932853d 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -219,3 +219,7 @@ def test_SphereAggregation_4D_agg_time() -> None: 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) -- 2.52.0 From a083bd30069b4b60222b4522b8dce1c68c9f2855 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 28 Mar 2023 06:53:05 +0200 Subject: [PATCH 5/7] requested changes --- docs/changes/latest.inc | 2 +- junifer/markers/parcel_aggregation.py | 4 ++-- junifer/markers/sphere_aggregation.py | 4 ++-- junifer/stats.py | 15 +++++++++------ 4 files changed, 14 insertions(+), 11 deletions(-) diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index ae7ad2827..9eef40a05 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -73,7 +73,7 @@ Enhancements - Add func:`junifer.stats.select` as an aggregation function that allows to select a subset of elements (:gh:`204` by `Fede Raimondo`_). -- Added ``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 spheres/parcels (: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 bfab9bd5b..c8c3f892c 100644 --- a/junifer/markers/parcel_aggregation.py +++ b/junifer/markers/parcel_aggregation.py @@ -34,8 +34,8 @@ class ParcelAggregation(BaseMarker): :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 `method` (only applicable to BOLD data). If None, it - will not operate on the time dimension (default None). + 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 diff --git a/junifer/markers/sphere_aggregation.py b/junifer/markers/sphere_aggregation.py index 9034142f9..1a30bccdb 100644 --- a/junifer/markers/sphere_aggregation.py +++ b/junifer/markers/sphere_aggregation.py @@ -39,8 +39,8 @@ class SphereAggregation(BaseMarker): 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 `method` (only applicable to BOLD data). If None, it - will not operate on the time dimension (default None). + 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 diff --git a/junifer/stats.py b/junifer/stats.py index 5318671d1..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 @@ -157,7 +157,10 @@ def winsorized_mean( def select( - data: np.ndarray, axis: int = 0, pick=None, drop=None + data: np.ndarray, + axis: int = 0, + pick: Optional[List[int]] = None, + drop: Optional[List[int]] = None, ) -> np.ndarray: """Select a subset of the data. @@ -167,9 +170,9 @@ def select( Data to select a subset from. axis : int, optional The axis to select a subset from (default 0). - pick : list, optional + pick : list of int, optional List of indices to select (default None). - drop : list, optional + drop : list of int, optional List of indices to drop (default None). Returns @@ -186,6 +189,6 @@ def select( 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) - out = data.take(pick, axis=axis) + pick = np.array(pick) # type: ignore + out = data.take(pick, axis=axis) # type: ignore return out -- 2.52.0 From ed0a79f6cc0b228079a37114c21e5f24cfa737bc Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 28 Mar 2023 07:01:47 +0200 Subject: [PATCH 6/7] Increase coverage --- junifer/markers/tests/test_parcel_aggregation.py | 10 ++++++++++ junifer/markers/tests/test_sphere_aggregation.py | 11 +++++++++++ junifer/tests/test_stats.py | 3 +++ 3 files changed, 24 insertions(+) diff --git a/junifer/markers/tests/test_parcel_aggregation.py b/junifer/markers/tests/test_parcel_aggregation.py index cbfe138f9..84d5d4c58 100644 --- a/junifer/markers/tests/test_parcel_aggregation.py +++ b/junifer/markers/tests/test_parcel_aggregation.py @@ -609,6 +609,16 @@ def test_ParcelAggregation_4D_agg_time(): 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 f4932853d..68ea96714 100644 --- a/junifer/markers/tests/test_sphere_aggregation.py +++ b/junifer/markers/tests/test_sphere_aggregation.py @@ -220,6 +220,17 @@ def test_SphereAggregation_4D_agg_time() -> None: 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/tests/test_stats.py b/junifer/tests/test_stats.py index 637d10984..bfddba9bc 100644 --- a/junifer/tests/test_stats.py +++ b/junifer/tests/test_stats.py @@ -125,3 +125,6 @@ def test_select() -> None: 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) -- 2.52.0 From 6151f2919eb3005a37d760bc2e5930921c1a3468 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 28 Mar 2023 08:38:52 +0200 Subject: [PATCH 7/7] even more coverage --- junifer/tests/test_stats.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py index bfddba9bc..c8b5fce63 100644 --- a/junifer/tests/test_stats.py +++ b/junifer/tests/test_stats.py @@ -128,3 +128,6 @@ def test_select() -> None: 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) -- 2.52.0