From d31cacf7581045617731ad7762aea718b5255e3e Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 13:40:13 +0100 Subject: [PATCH 1/9] Add allow_overlap and allow empty spheres in sphere aggregation --- .../nilearn/junifer_nifti_spheres_masker.py | 158 +++++++++++++++++- .../test_junifer_nifti_spheres_masker.py | 5 +- junifer/markers/falff/falff_spheres.py | 6 + .../edge_functional_connectivity_spheres.py | 6 + .../functional_connectivity_spheres.py | 6 + junifer/markers/reho/reho_spheres.py | 6 + junifer/markers/sphere_aggregation.py | 6 + .../temporal_snr/temporal_snr_spheres.py | 6 + 8 files changed, 194 insertions(+), 5 deletions(-) diff --git a/junifer/external/nilearn/junifer_nifti_spheres_masker.py b/junifer/external/nilearn/junifer_nifti_spheres_masker.py index 8787925e0..fe6ad40c1 100644 --- a/junifer/external/nilearn/junifer_nifti_spheres_masker.py +++ b/junifer/external/nilearn/junifer_nifti_spheres_masker.py @@ -6,14 +6,19 @@ from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union import numpy as np +from nilearn import image, masking from nilearn._utils.class_inspect import get_params from nilearn._utils.niimg import img_data_dtype -from nilearn._utils.niimg_conversions import check_niimg_4d +from nilearn._utils.niimg_conversions import ( + check_niimg_4d, + check_niimg_3d, + _safe_get_data, +) from nilearn.maskers import NiftiSpheresMasker from nilearn.maskers.base_masker import _filter_and_extract -from nilearn.maskers.nifti_spheres_masker import _iter_signals_from_spheres +from sklearn import neighbors -from ...utils import raise_error +from ...utils import raise_error, warn_with_log if TYPE_CHECKING: @@ -59,6 +64,153 @@ DAMAGE. """ +def _apply_mask_and_get_affinity(seeds, niimg, radius, allow_overlap, + mask_img=None): + """Apply mask and get affinity matrix for given seeds. + + Utility function to get only the rows which are occupied by sphere at + given seed locations and the provided radius. Rows are in target_affine and + target_shape space. + + Parameters + ---------- + seeds : List of triplets of coordinates in native space + Seed definitions. List of coordinates of the seeds in the same space + as target_affine. + niimg : 3D/4D Niimg-like object + See :ref:`extracting_data`. + Images to process. + If a 3D niimg is provided, a singleton dimension will be added to + the output to represent the single scan in the niimg. + radius : float + Indicates, in millimeters, the radius for the sphere around the seed. + allow_overlap : boolean + If False, a ValueError is raised if VOIs overlap + mask_img : Niimg-like object, optional + Mask to apply to regions before extracting signals. If niimg is None, + mask_img is used as a reference space in which the spheres 'indices are + placed. + Returns + ------- + X : 2D numpy.ndarray + Signal for each brain voxel in the (masked) niimgs. + shape: (number of scans, number of voxels) + A : scipy.sparse.lil_matrix + Contains the boolean indices for each sphere. + shape: (number of seeds, number of voxels) + """ + seeds = list(seeds) + + # Compute world coordinates of all in-mask voxels. + if niimg is None: + mask, affine = masking._load_mask_img(mask_img) + # Get coordinate for all voxels inside of mask + mask_coords = np.asarray(np.nonzero(mask)).T.tolist() + X = None + + elif mask_img is not None: + affine = niimg.affine + mask_img = check_niimg_3d(mask_img) + mask_img = image.resample_img( + mask_img, + target_affine=affine, + target_shape=niimg.shape[:3], + interpolation='nearest', + ) + mask, _ = masking._load_mask_img(mask_img) + mask_coords = list(zip(*np.where(mask != 0))) + + X = masking._apply_mask_fmri(niimg, mask_img) + + elif niimg is not None: + affine = niimg.affine + if np.isnan(np.sum(_safe_get_data(niimg))): + warn_with_log( + 'The imgs you have fed into fit_transform() contains NaN ' + 'values which will be converted to zeroes.' + ) + X = _safe_get_data(niimg, True).reshape([-1, niimg.shape[3]]).T + else: + X = _safe_get_data(niimg).reshape([-1, niimg.shape[3]]).T + + mask_coords = list(np.ndindex(niimg.shape[:3])) + + else: + raise_error("Either a niimg or a mask_img must be provided.") + + # For each seed, get coordinates of nearest voxel + nearests = [] + for sx, sy, sz in seeds: + nearest = np.round(image.resampling.coord_transform( + sx, sy, sz, np.linalg.inv(affine) + )) + nearest = nearest.astype(int) + nearest = (nearest[0], nearest[1], nearest[2]) + try: + nearests.append(mask_coords.index(nearest)) + except ValueError: + nearests.append(None) + + mask_coords = np.asarray(list(zip(*mask_coords))) + mask_coords = image.resampling.coord_transform( + mask_coords[0], mask_coords[1], mask_coords[2], affine + ) + mask_coords = np.asarray(mask_coords).T + + clf = neighbors.NearestNeighbors(radius=radius) + A = clf.fit(mask_coords).radius_neighbors_graph(seeds) + A = A.tolil() + for i, nearest in enumerate(nearests): + if nearest is None: + continue + + A[i, nearest] = True + + # Include the voxel containing the seed itself if not masked + mask_coords = mask_coords.astype(int).tolist() + for i, seed in enumerate(seeds): + try: + A[i, mask_coords.index(list(map(int, seed)))] = True + except ValueError: + # seed is not in the mask + pass + + if (not allow_overlap) and np.any(A.sum(axis=0) >= 2): + raise_error('Overlap detected between spheres') + + return X, A + + +def _iter_signals_from_spheres(seeds, niimg, radius, allow_overlap, + mask_img=None): + """Iterate over spheres. + + Parameters + ---------- + seeds : :obj:`list` of triplets of coordinates in native space + Seed definitions. List of coordinates of the seeds in the same space + as the images (typically MNI or TAL). + niimg : 3D/4D Niimg-like object + See :ref:`extracting_data`. + Images to process. + If a 3D niimg is provided, a singleton dimension will be added to + the output to represent the single scan in the niimg. + radius: float + Indicates, in millimeters, the radius for the sphere around the seed. + allow_overlap: boolean + If False, an error is raised if the maps overlaps (ie at least two + maps have a non-zero value for the same voxel). + mask_img : Niimg-like object, optional + See :ref:`extracting_data`. + Mask to apply to regions before extracting signals. + """ + X, A = _apply_mask_and_get_affinity(seeds, niimg, radius, + allow_overlap, + mask_img=mask_img) + for _, row in enumerate(A.rows): + yield X[:, row] + + class _JuniferExtractionFunctor: """Functor to extract signals from spheres. diff --git a/junifer/external/nilearn/tests/test_junifer_nifti_spheres_masker.py b/junifer/external/nilearn/tests/test_junifer_nifti_spheres_masker.py index f771d3db5..b87ad8b61 100644 --- a/junifer/external/nilearn/tests/test_junifer_nifti_spheres_masker.py +++ b/junifer/external/nilearn/tests/test_junifer_nifti_spheres_masker.py @@ -216,8 +216,9 @@ def test_small_radius() -> None: radius=0.1, mask_img=nibabel.Nifti1Image(mask, affine), ) - with pytest.raises(ValueError, match="These spheres are empty"): - masker.fit_transform(nibabel.Nifti1Image(data, affine)) + + out = masker.fit_transform(nibabel.Nifti1Image(data, affine)) + assert np.isnan(out).all() masker = JuniferNiftiSpheresMasker( seeds=[seed], diff --git a/junifer/markers/falff/falff_spheres.py b/junifer/markers/falff/falff_spheres.py index 892816ac3..574fa66f3 100644 --- a/junifer/markers/falff/falff_spheres.py +++ b/junifer/markers/falff/falff_spheres.py @@ -27,6 +27,9 @@ class AmplitudeLowFrequencyFluctuationSpheres( The radius of the sphere in mm. If None, the signal will be extracted from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` for more information (default None). + allow_overlap : bool, optional + Whether to allow overlapping spheres. If False, an error is raised if + the spheres overlap (default is False). fractional : bool Whether to compute fractional ALFF. highpass : positive float, optional @@ -72,6 +75,7 @@ class AmplitudeLowFrequencyFluctuationSpheres( coords: str, fractional: bool, radius: Optional[float] = None, + allow_overlap: bool = False, highpass: float = 0.01, lowpass: float = 0.1, tr: Optional[float] = None, @@ -83,6 +87,7 @@ class AmplitudeLowFrequencyFluctuationSpheres( ) -> None: self.coords = coords self.radius = radius + self.allow_overlap = allow_overlap self.masks = masks self.method = method self.method_params = method_params @@ -124,6 +129,7 @@ class AmplitudeLowFrequencyFluctuationSpheres( pa = SphereAggregation( coords=self.coords, radius=self.radius, + allow_overlap=self.allow_overlap, method=self.method, method_params=self.method_params, masks=self.masks, diff --git a/junifer/markers/functional_connectivity/edge_functional_connectivity_spheres.py b/junifer/markers/functional_connectivity/edge_functional_connectivity_spheres.py index adcd065dc..2170f6ac9 100644 --- a/junifer/markers/functional_connectivity/edge_functional_connectivity_spheres.py +++ b/junifer/markers/functional_connectivity/edge_functional_connectivity_spheres.py @@ -25,6 +25,9 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase): The radius of the sphere in mm. If None, the signal will be extracted from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` for more information (default None). + allow_overlap : bool, optional + Whether to allow overlapping spheres. If False, an error is raised if + the spheres overlap (default is False). agg_method : str, optional The aggregation method to use. See :func:`junifer.stats.get_aggfunc_by_name` for more information @@ -59,6 +62,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase): self, coords: str, radius: Optional[float] = None, + allow_overlap: bool = False, agg_method: str = "mean", agg_method_params: Optional[Dict] = None, cor_method: str = "covariance", @@ -68,6 +72,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase): ) -> None: self.coords = coords self.radius = radius + self.allow_overlap = allow_overlap if radius is None or radius <= 0: raise_error(f"radius should be > 0: provided {radius}") super().__init__( @@ -84,6 +89,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase): sphere_aggregation = SphereAggregation( coords=self.coords, radius=self.radius, + allow_overlap=self.allow_overlap, method=self.agg_method, method_params=self.agg_method_params, masks=self.masks, diff --git a/junifer/markers/functional_connectivity/functional_connectivity_spheres.py b/junifer/markers/functional_connectivity/functional_connectivity_spheres.py index f6552c81e..4fa5a837e 100644 --- a/junifer/markers/functional_connectivity/functional_connectivity_spheres.py +++ b/junifer/markers/functional_connectivity/functional_connectivity_spheres.py @@ -26,6 +26,9 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase): The radius of the sphere in mm. If None, the signal will be extracted from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` for more information (default None). + allow_overlap : bool, optional + Whether to allow overlapping spheres. If False, an error is raised if + the spheres overlap (default is False). agg_method : str, optional The aggregation method to use. See :func:`junifer.stats.get_aggfunc_by_name` for more information @@ -53,6 +56,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase): self, coords: str, radius: Optional[float] = None, + allow_overlap: bool = False, agg_method: str = "mean", agg_method_params: Optional[Dict] = None, cor_method: str = "covariance", @@ -62,6 +66,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase): ) -> None: self.coords = coords self.radius = radius + self.allow_overlap = allow_overlap if radius is None or radius <= 0: raise_error(f"radius should be > 0: provided {radius}") super().__init__( @@ -78,6 +83,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase): sphere_aggregation = SphereAggregation( coords=self.coords, radius=self.radius, + allow_overlap=self.allow_overlap, method=self.agg_method, method_params=self.agg_method_params, masks=self.masks, diff --git a/junifer/markers/reho/reho_spheres.py b/junifer/markers/reho/reho_spheres.py index 3049cc4ad..4db4ed928 100644 --- a/junifer/markers/reho/reho_spheres.py +++ b/junifer/markers/reho/reho_spheres.py @@ -28,6 +28,9 @@ class ReHoSpheres(ReHoBase): extracted from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` for more information (default None). + allow_overlap : bool, optional + Whether to allow overlapping spheres. If False, an error is raised if + the spheres overlap (default is False). use_afni : bool, optional Whether to use AFNI for computing. If None, will use AFNI only if available (default None). @@ -93,6 +96,7 @@ class ReHoSpheres(ReHoBase): self, coords: str, radius: Optional[float] = None, + allow_overlap: bool = False, use_afni: Optional[bool] = None, reho_params: Optional[Dict] = None, agg_method: str = "mean", @@ -102,6 +106,7 @@ class ReHoSpheres(ReHoBase): ) -> None: self.coords = coords self.radius = radius + self.allow_overlap = allow_overlap self.reho_params = reho_params self.agg_method = agg_method self.agg_method_params = agg_method_params @@ -142,6 +147,7 @@ class ReHoSpheres(ReHoBase): sphere_aggregation = SphereAggregation( coords=self.coords, radius=self.radius, + allow_overlap=self.allow_overlap, method=self.agg_method, method_params=self.agg_method_params, masks=self.masks, diff --git a/junifer/markers/sphere_aggregation.py b/junifer/markers/sphere_aggregation.py index ae90cef58..756e1d189 100644 --- a/junifer/markers/sphere_aggregation.py +++ b/junifer/markers/sphere_aggregation.py @@ -28,6 +28,9 @@ class SphereAggregation(BaseMarker): extracted from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` for more information (default None). + allow_overlap : bool, optional + Whether to allow overlapping spheres. If False, an error is raised if + the spheres overlap (default is False). method : str, optional The aggregation method to use. See :func:`junifer.stats.get_aggfunc_by_name` for more information @@ -54,6 +57,7 @@ class SphereAggregation(BaseMarker): self, coords: str, radius: Optional[float] = None, + allow_overlap: bool = False, method: str = "mean", method_params: Optional[Dict[str, Any]] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None, @@ -62,6 +66,7 @@ class SphereAggregation(BaseMarker): ) -> None: self.coords = coords self.radius = radius + self.allow_overlap = allow_overlap self.method = method self.method_params = method_params or {} self.masks = masks @@ -147,6 +152,7 @@ class SphereAggregation(BaseMarker): masker = JuniferNiftiSpheresMasker( seeds=coords, radius=self.radius, + allow_overlap=self.allow_overlap, mask_img=mask_img, agg_func=agg_func, ) diff --git a/junifer/markers/temporal_snr/temporal_snr_spheres.py b/junifer/markers/temporal_snr/temporal_snr_spheres.py index 402d4aa47..712a1b713 100644 --- a/junifer/markers/temporal_snr/temporal_snr_spheres.py +++ b/junifer/markers/temporal_snr/temporal_snr_spheres.py @@ -24,6 +24,9 @@ class TemporalSNRSpheres(TemporalSNRBase): The radius of the sphere in mm. If None, the signal will be extracted from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` for more information (default None). + allow_overlap : bool, optional + Whether to allow overlapping spheres. If False, an error is raised if + the spheres overlap (default is False). agg_method : str, optional The aggregation method to use. See :func:`junifer.stats.get_aggfunc_by_name` for more information @@ -45,6 +48,7 @@ class TemporalSNRSpheres(TemporalSNRBase): self, coords: str, radius: Optional[float] = None, + allow_overlap: bool = False, agg_method: str = "mean", agg_method_params: Optional[Dict] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None, @@ -52,6 +56,7 @@ class TemporalSNRSpheres(TemporalSNRBase): ) -> None: self.coords = coords self.radius = radius + self.allow_overlap = allow_overlap if radius is None or radius <= 0: raise_error(f"radius should be > 0: provided {radius}") super().__init__( @@ -91,6 +96,7 @@ class TemporalSNRSpheres(TemporalSNRBase): sphere_aggregation = SphereAggregation( coords=self.coords, radius=self.radius, + allow_overlap=self.allow_overlap, method=self.agg_method, method_params=self.agg_method_params, masks=self.masks, -- 2.52.0 From b4ed1bfa8e6ff9f94f6f670699b2bc99818db8f8 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 13:47:53 +0100 Subject: [PATCH 2/9] add changes --- docs/changes/latest.inc | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index c7df47c88..762816efc 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -51,6 +51,10 @@ Enhancements - Add ``pre_run`` parameter to ``_queue_condor`` (:gh:`188` by `Fede Raimondo`_). +- Allow for empty spheres in :class:`junifer.external.nilearn.JuniferNiftiSpheresMasker`, that will result in NaNs (:gh:`190` by `Fede Raimondo`_). + +- Expose `allow_overlap` parameter in :class:`junifer.markers.SphereAgregation` and related markers (:gh:`190` by `Fede Raimondo`_). + Bugs ~~~~ -- 2.52.0 From 81dfb34402c16f05bd5de038f4a5a6c2d5358a99 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 13:48:19 +0100 Subject: [PATCH 3/9] linter --- .../nilearn/junifer_nifti_spheres_masker.py | 32 ++++++++++--------- 1 file changed, 17 insertions(+), 15 deletions(-) diff --git a/junifer/external/nilearn/junifer_nifti_spheres_masker.py b/junifer/external/nilearn/junifer_nifti_spheres_masker.py index fe6ad40c1..a6de1eb8d 100644 --- a/junifer/external/nilearn/junifer_nifti_spheres_masker.py +++ b/junifer/external/nilearn/junifer_nifti_spheres_masker.py @@ -64,10 +64,11 @@ DAMAGE. """ -def _apply_mask_and_get_affinity(seeds, niimg, radius, allow_overlap, - mask_img=None): +def _apply_mask_and_get_affinity( + seeds, niimg, radius, allow_overlap, mask_img=None +): """Apply mask and get affinity matrix for given seeds. - + Utility function to get only the rows which are occupied by sphere at given seed locations and the provided radius. Rows are in target_affine and target_shape space. @@ -115,7 +116,7 @@ def _apply_mask_and_get_affinity(seeds, niimg, radius, allow_overlap, mask_img, target_affine=affine, target_shape=niimg.shape[:3], - interpolation='nearest', + interpolation="nearest", ) mask, _ = masking._load_mask_img(mask_img) mask_coords = list(zip(*np.where(mask != 0))) @@ -126,8 +127,8 @@ def _apply_mask_and_get_affinity(seeds, niimg, radius, allow_overlap, affine = niimg.affine if np.isnan(np.sum(_safe_get_data(niimg))): warn_with_log( - 'The imgs you have fed into fit_transform() contains NaN ' - 'values which will be converted to zeroes.' + "The imgs you have fed into fit_transform() contains NaN " + "values which will be converted to zeroes." ) X = _safe_get_data(niimg, True).reshape([-1, niimg.shape[3]]).T else: @@ -141,9 +142,9 @@ def _apply_mask_and_get_affinity(seeds, niimg, radius, allow_overlap, # For each seed, get coordinates of nearest voxel nearests = [] for sx, sy, sz in seeds: - nearest = np.round(image.resampling.coord_transform( - sx, sy, sz, np.linalg.inv(affine) - )) + nearest = np.round( + image.resampling.coord_transform(sx, sy, sz, np.linalg.inv(affine)) + ) nearest = nearest.astype(int) nearest = (nearest[0], nearest[1], nearest[2]) try: @@ -176,13 +177,14 @@ def _apply_mask_and_get_affinity(seeds, niimg, radius, allow_overlap, pass if (not allow_overlap) and np.any(A.sum(axis=0) >= 2): - raise_error('Overlap detected between spheres') + raise_error("Overlap detected between spheres") return X, A -def _iter_signals_from_spheres(seeds, niimg, radius, allow_overlap, - mask_img=None): +def _iter_signals_from_spheres( + seeds, niimg, radius, allow_overlap, mask_img=None +): """Iterate over spheres. Parameters @@ -204,9 +206,9 @@ def _iter_signals_from_spheres(seeds, niimg, radius, allow_overlap, See :ref:`extracting_data`. Mask to apply to regions before extracting signals. """ - X, A = _apply_mask_and_get_affinity(seeds, niimg, radius, - allow_overlap, - mask_img=mask_img) + X, A = _apply_mask_and_get_affinity( + seeds, niimg, radius, allow_overlap, mask_img=mask_img + ) for _, row in enumerate(A.rows): yield X[:, row] -- 2.52.0 From 7ec4711eae5ff5ddf37c43f70616fd3a3193ef50 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 13:56:04 +0100 Subject: [PATCH 4/9] Add count aggregation function --- docs/changes/latest.inc | 2 ++ junifer/stats.py | 29 ++++++++++++++++++++++++++++- junifer/tests/test_stats.py | 16 +++++++++++++++- 3 files changed, 45 insertions(+), 2 deletions(-) diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index 762816efc..b3346cc60 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -55,6 +55,8 @@ Enhancements - Expose `allow_overlap` parameter in :class:`junifer.markers.SphereAgregation` and related markers (:gh:`190` by `Fede Raimondo`_). +- Add aggregation function :func:`junifer.stats.count` that returns the number of elements in a given axis. This allows to count the number of voxels per sphere/parcel when used as `agg_func` in markers (:gh:`190` by `Fede Raimondo`_). + Bugs ~~~~ diff --git a/junifer/stats.py b/junifer/stats.py index 7cc369ae4..6679a16c6 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -28,6 +28,7 @@ def get_aggfunc_by_name( * ``mean`` -> :func:`numpy.mean` * ``std`` -> :func:`numpy.std` * ``trim_mean`` -> :func:`scipy.stats.trim_mean` + * ``count`` -> :func:`junifer.stats.count` func_params : dict, optional Parameters to pass to the function. @@ -42,7 +43,13 @@ def get_aggfunc_by_name( from functools import partial # local import to avoid sphinx error # check validity of names - _valid_func_names = {"winsorized_mean", "mean", "std", "trim_mean"} + _valid_func_names = { + "winsorized_mean", + "mean", + "std", + "trim_mean", + "count", + } if func_params is None: func_params = {} # apply functions @@ -74,6 +81,8 @@ def get_aggfunc_by_name( func = np.std elif name == "trim_mean": func = partial(trim_mean, **func_params) + elif name == "count": + func = count else: raise_error( f"Function {name} unknown. Please provide any of " @@ -82,6 +91,24 @@ def get_aggfunc_by_name( return func +def count(data: np.ndarray, axis: int = 0) -> np.ndarray: + """Count the number elements along the given axis. + + Parameters + ---------- + data : numpy.ndarray + Data to count elements on. + axis : int, optional + The axis to count elements on (default 0). + + Returns + ------- + numpy.ndarray + Number of lements along the given axis. + """ + return data.shape[axis] + + def winsorized_mean( data: np.ndarray, axis: Optional[int] = None, **win_params ) -> np.ndarray: diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py index ef36b1071..7ef6b6387 100644 --- a/junifer/tests/test_stats.py +++ b/junifer/tests/test_stats.py @@ -8,7 +8,7 @@ from typing import Dict, Optional import numpy as np import pytest -from junifer.stats import get_aggfunc_by_name, winsorized_mean +from junifer.stats import get_aggfunc_by_name, winsorized_mean, count @pytest.mark.parametrize( @@ -17,6 +17,7 @@ from junifer.stats import get_aggfunc_by_name, winsorized_mean ("winsorized_mean", {"limits": [0.2, 0.7]}), ("mean", None), ("std", None), + ("count", None), ("trim_mean", None), ("trim_mean", {"proportiontocut": 0.1}), ], @@ -74,3 +75,16 @@ def test_winsorized_mean() -> None: input = np.array([22, 4, 9, 8, 5, 3, 7, 2, 1, 6]) output = winsorized_mean(input, limits=[0.1, 0.1]) assert output == 5.5 + + +def test_count() -> None: + """Test count.""" + input = np.zeros((10, 3)) + assert count(input, axis=-1) == 3 + assert count(input, axis=1) == 3 + assert count(input, axis=0) == 10 + + input = np.zeros((10, 0)) + assert count(input, axis=-1) == 0 + assert count(input, axis=1) == 0 + assert count(input, axis=0) == 10 -- 2.52.0 From 923eddb55f9bda0ca2db7796e36aaa712f857db2 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 13:56:27 +0100 Subject: [PATCH 5/9] change type --- junifer/stats.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/stats.py b/junifer/stats.py index 6679a16c6..f12d4fe3e 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -91,7 +91,7 @@ def get_aggfunc_by_name( return func -def count(data: np.ndarray, axis: int = 0) -> np.ndarray: +def count(data: np.ndarray, axis: int = 0) -> int: """Count the number elements along the given axis. Parameters -- 2.52.0 From 00764dc12e5796b9b6c5824931ba8f687a7ce915 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 15:00:24 +0100 Subject: [PATCH 6/9] fix docs --- docs/api/index.rst | 10 ++++++++++ docs/api/nilearn.rst | 8 ++++++++ docs/changes/latest.inc | 4 ++-- .../nilearn/junifer_nifti_spheres_masker.py | 17 ++++++++++++----- 4 files changed, 32 insertions(+), 7 deletions(-) create mode 100644 docs/api/nilearn.rst diff --git a/docs/api/index.rst b/docs/api/index.rst index 4abc5afa7..875204997 100644 --- a/docs/api/index.rst +++ b/docs/api/index.rst @@ -36,3 +36,13 @@ Configs :caption: Contents: configs + + +External +-------- + +.. toctree:: + :maxdepth: 2 + :caption: Contents: + + nilearn.rst \ No newline at end of file diff --git a/docs/api/nilearn.rst b/docs/api/nilearn.rst new file mode 100644 index 000000000..100ef9fb5 --- /dev/null +++ b/docs/api/nilearn.rst @@ -0,0 +1,8 @@ +Nilearn +======= + +This package provides re-implementations of some of the functions in `nilearn`_. + +.. automodule:: junifer.external.nilearn + :members: + :imported-members: diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index b3346cc60..875fcfa3f 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -53,9 +53,9 @@ Enhancements - Allow for empty spheres in :class:`junifer.external.nilearn.JuniferNiftiSpheresMasker`, that will result in NaNs (:gh:`190` by `Fede Raimondo`_). -- Expose `allow_overlap` parameter in :class:`junifer.markers.SphereAgregation` and related markers (:gh:`190` by `Fede Raimondo`_). +- Expose ``allow_overlap`` parameter in :class:`junifer.markers.SphereAggregation` and related markers (:gh:`190` by `Fede Raimondo`_). -- Add aggregation function :func:`junifer.stats.count` that returns the number of elements in a given axis. This allows to count the number of voxels per sphere/parcel when used as `agg_func` in markers (:gh:`190` by `Fede Raimondo`_). +- Add aggregation function :func:`junifer.stats.count` that returns the number of elements in a given axis. This allows to count the number of voxels per sphere/parcel when used as ``method`` in markers (:gh:`190` by `Fede Raimondo`_). Bugs ~~~~ diff --git a/junifer/external/nilearn/junifer_nifti_spheres_masker.py b/junifer/external/nilearn/junifer_nifti_spheres_masker.py index a6de1eb8d..39e39bd02 100644 --- a/junifer/external/nilearn/junifer_nifti_spheres_masker.py +++ b/junifer/external/nilearn/junifer_nifti_spheres_masker.py @@ -303,9 +303,16 @@ class _JuniferExtractionFunctor: class JuniferNiftiSpheresMasker(NiftiSpheresMasker): """Class for custom NiftiSpheresMasker. + Differs from :class:`nilearn.maskers.NiftiSpheresMasker` in the following + ways: + + * it allows to pass any callable as the ``agg_func`` parameter. + * empty spheres do not create an error. Insted, ``agg_func`` is applied to + an empty array and the result is passed. + Parameters ---------- - seeds : list of triplet of coordinates in native space + seeds : list of float Seed definitions. List of coordinates of the seeds in the same space as the images (typically MNI or TAL). radius : float, optional @@ -317,13 +324,13 @@ class JuniferNiftiSpheresMasker(NiftiSpheresMasker): The function to aggregate signals using (default numpy.mean). allow_overlap : bool, optional If False, an error is raised if the maps overlap (default None). - dtype : any type that can be coerced into a numpy dtype or "auto", optional + dtype : numpy.dtype or "auto", optional The dtype for the extraction. If "auto", the data will be converted to int32 if dtype is discrete and float32 if it is continuous (default None). **kwargs Keyword arguments are passed to the - :func:`nilearn.maskers.NiftiSpheresMasker`. + :class:`nilearn.maskers.NiftiSpheresMasker`. """ @@ -362,11 +369,11 @@ class JuniferNiftiSpheresMasker(NiftiSpheresMasker): Images to process. If a 3D niimg is provided, a singleton dimension will be added to the output to represent the single scan in the niimg. - confounds : CSV file or array-like or pandas.DataFrame, optional + confounds : pandas.DataFrame, optional This parameter is passed to :func:`nilearn.signal.clean`. Please see the related documentation for details. shape: (number of scans, number of confounds) - sample_mask : Any type compatible with numpy-array indexing, optional + sample_mask : np.ndarray, list or tuple, optional Masks the niimgs along time/fourth dimension to perform scrubbing (remove volumes with high motion) and/or non-steady-state volumes. This parameter is passed to :func:`nilearn.signal.clean`. -- 2.52.0 From 0bc39983eb4e3c49bb1be6aab40878926a5a58d3 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 15:04:35 +0100 Subject: [PATCH 7/9] codespell --- junifer/external/nilearn/junifer_nifti_spheres_masker.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/external/nilearn/junifer_nifti_spheres_masker.py b/junifer/external/nilearn/junifer_nifti_spheres_masker.py index 39e39bd02..f301a58fe 100644 --- a/junifer/external/nilearn/junifer_nifti_spheres_masker.py +++ b/junifer/external/nilearn/junifer_nifti_spheres_masker.py @@ -307,7 +307,7 @@ class JuniferNiftiSpheresMasker(NiftiSpheresMasker): ways: * it allows to pass any callable as the ``agg_func`` parameter. - * empty spheres do not create an error. Insted, ``agg_func`` is applied to + * empty spheres do not create an error. Instead, ``agg_func`` is applied to an empty array and the result is passed. Parameters -- 2.52.0 From 4e026b7d6cb5b77e81662e5540436d967a376c2c Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 15 Mar 2023 15:51:38 +0100 Subject: [PATCH 8/9] fix comments --- junifer/stats.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/stats.py b/junifer/stats.py index f12d4fe3e..a26a8142b 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -103,8 +103,8 @@ def count(data: np.ndarray, axis: int = 0) -> int: Returns ------- - numpy.ndarray - Number of lements along the given axis. + int + Number of elements along the given axis. """ return data.shape[axis] -- 2.52.0 From f312ae5e713ae5844e408705449d38dee5611bdb Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 15 Mar 2023 15:59:08 +0100 Subject: [PATCH 9/9] chore: isort --- junifer/external/nilearn/junifer_nifti_spheres_masker.py | 4 ++-- junifer/tests/test_stats.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/junifer/external/nilearn/junifer_nifti_spheres_masker.py b/junifer/external/nilearn/junifer_nifti_spheres_masker.py index f301a58fe..3f2517515 100644 --- a/junifer/external/nilearn/junifer_nifti_spheres_masker.py +++ b/junifer/external/nilearn/junifer_nifti_spheres_masker.py @@ -10,9 +10,9 @@ from nilearn import image, masking from nilearn._utils.class_inspect import get_params from nilearn._utils.niimg import img_data_dtype from nilearn._utils.niimg_conversions import ( - check_niimg_4d, - check_niimg_3d, _safe_get_data, + check_niimg_3d, + check_niimg_4d, ) from nilearn.maskers import NiftiSpheresMasker from nilearn.maskers.base_masker import _filter_and_extract diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py index 7ef6b6387..37c50114e 100644 --- a/junifer/tests/test_stats.py +++ b/junifer/tests/test_stats.py @@ -8,7 +8,7 @@ from typing import Dict, Optional import numpy as np import pytest -from junifer.stats import get_aggfunc_by_name, winsorized_mean, count +from junifer.stats import count, get_aggfunc_by_name, winsorized_mean @pytest.mark.parametrize( -- 2.52.0