Add allow_overlap and allow empty spheres in sphere aggregation #190

Merged
fraimondo merged 9 commits from enh/empty_spheres into main 2023-03-15 15:34:17 +00:00
13 changed files with 275 additions and 12 deletions

View file

@ -36,3 +36,13 @@ Configs
:caption: Contents: :caption: Contents:
configs configs
External
--------
.. toctree::
:maxdepth: 2
:caption: Contents:
nilearn.rst

8
docs/api/nilearn.rst Normal file
View file

@ -0,0 +1,8 @@
Nilearn
=======
This package provides re-implementations of some of the functions in `nilearn`_.
.. automodule:: junifer.external.nilearn
:members:
:imported-members:

View file

@ -51,6 +51,12 @@ Enhancements
synchon commented 2023-03-15 13:16:38 +00:00 (Migrated from github.com)
``allow_overlap``
``` ``allow_overlap`` ```
- Add ``pre_run`` parameter to ``_queue_condor`` (:gh:`188` by `Fede Raimondo`_). - 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.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 ``method`` in markers (:gh:`190` by `Fede Raimondo`_).
Bugs Bugs
~~~~ ~~~~

View file

@ -6,14 +6,19 @@
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
import numpy as np import numpy as np
from nilearn import image, masking
from nilearn._utils.class_inspect import get_params from nilearn._utils.class_inspect import get_params
from nilearn._utils.niimg import img_data_dtype from nilearn._utils.niimg import img_data_dtype
from nilearn._utils.niimg_conversions import check_niimg_4d from nilearn._utils.niimg_conversions import (
_safe_get_data,
check_niimg_3d,
check_niimg_4d,
)
from nilearn.maskers import NiftiSpheresMasker from nilearn.maskers import NiftiSpheresMasker
from nilearn.maskers.base_masker import _filter_and_extract 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: if TYPE_CHECKING:
@ -59,6 +64,155 @@ 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
synchon commented 2023-03-15 13:18:55 +00:00 (Migrated from github.com)

This is for adding the custom check of empty spheres right?

This is for adding the custom check of empty spheres right?
fraimondo commented 2023-03-15 14:51:55 +00:00 (Migrated from github.com)

Indeed the only change from nilearn was to remove the check.

Indeed the only change from nilearn was to remove the check.
synchon commented 2023-03-15 14:56:04 +00:00 (Migrated from github.com)

Sounds good.

Sounds good.
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: class _JuniferExtractionFunctor:
"""Functor to extract signals from spheres. """Functor to extract signals from spheres.
@ -149,9 +303,16 @@ class _JuniferExtractionFunctor:
class JuniferNiftiSpheresMasker(NiftiSpheresMasker): class JuniferNiftiSpheresMasker(NiftiSpheresMasker):
"""Class for custom 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. Instead, ``agg_func`` is applied to
an empty array and the result is passed.
Parameters 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 Seed definitions. List of coordinates of the seeds in the same space
as the images (typically MNI or TAL). as the images (typically MNI or TAL).
radius : float, optional radius : float, optional
@ -163,13 +324,13 @@ class JuniferNiftiSpheresMasker(NiftiSpheresMasker):
The function to aggregate signals using (default numpy.mean). The function to aggregate signals using (default numpy.mean).
allow_overlap : bool, optional allow_overlap : bool, optional
If False, an error is raised if the maps overlap (default None). 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 The dtype for the extraction. If "auto", the data will be converted to
int32 if dtype is discrete and float32 if it is continuous int32 if dtype is discrete and float32 if it is continuous
(default None). (default None).
**kwargs **kwargs
Keyword arguments are passed to the Keyword arguments are passed to the
:func:`nilearn.maskers.NiftiSpheresMasker`. :class:`nilearn.maskers.NiftiSpheresMasker`.
""" """
@ -208,11 +369,11 @@ class JuniferNiftiSpheresMasker(NiftiSpheresMasker):
Images to process. Images to process.
If a 3D niimg is provided, a singleton dimension will be added to If a 3D niimg is provided, a singleton dimension will be added to
the output to represent the single scan in the niimg. 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`. This parameter is passed to :func:`nilearn.signal.clean`.
Please see the related documentation for details. Please see the related documentation for details.
shape: (number of scans, number of confounds) 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 Masks the niimgs along time/fourth dimension to perform scrubbing
(remove volumes with high motion) and/or non-steady-state volumes. (remove volumes with high motion) and/or non-steady-state volumes.
This parameter is passed to :func:`nilearn.signal.clean`. This parameter is passed to :func:`nilearn.signal.clean`.

View file

@ -216,8 +216,9 @@ def test_small_radius() -> None:
radius=0.1, radius=0.1,
mask_img=nibabel.Nifti1Image(mask, affine), 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( masker = JuniferNiftiSpheresMasker(
seeds=[seed], seeds=[seed],

View file

@ -27,6 +27,9 @@ class AmplitudeLowFrequencyFluctuationSpheres(
The radius of the sphere in mm. If None, the signal will be extracted The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information (default None). 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 fractional : bool
Whether to compute fractional ALFF. Whether to compute fractional ALFF.
highpass : positive float, optional highpass : positive float, optional
@ -72,6 +75,7 @@ class AmplitudeLowFrequencyFluctuationSpheres(
coords: str, coords: str,
fractional: bool, fractional: bool,
radius: Optional[float] = None, radius: Optional[float] = None,
allow_overlap: bool = False,
highpass: float = 0.01, highpass: float = 0.01,
lowpass: float = 0.1, lowpass: float = 0.1,
tr: Optional[float] = None, tr: Optional[float] = None,
@ -83,6 +87,7 @@ class AmplitudeLowFrequencyFluctuationSpheres(
) -> None: ) -> None:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.allow_overlap = allow_overlap
self.masks = masks self.masks = masks
self.method = method self.method = method
self.method_params = method_params self.method_params = method_params
@ -124,6 +129,7 @@ class AmplitudeLowFrequencyFluctuationSpheres(
pa = SphereAggregation( pa = SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.method, method=self.method,
method_params=self.method_params, method_params=self.method_params,
masks=self.masks, masks=self.masks,

View file

@ -25,6 +25,9 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
The radius of the sphere in mm. If None, the signal will be extracted The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information (default None). 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 agg_method : str, optional
The aggregation method to use. The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information See :func:`junifer.stats.get_aggfunc_by_name` for more information
@ -59,6 +62,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
self, self,
coords: str, coords: str,
radius: Optional[float] = None, radius: Optional[float] = None,
allow_overlap: bool = False,
agg_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance", cor_method: str = "covariance",
@ -68,6 +72,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
) -> None: ) -> None:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.allow_overlap = allow_overlap
if radius is None or radius <= 0: if radius is None or radius <= 0:
raise_error(f"radius should be > 0: provided {radius}") raise_error(f"radius should be > 0: provided {radius}")
super().__init__( super().__init__(
@ -84,6 +89,7 @@ class EdgeCentricFCSpheres(FunctionalConnectivityBase):
sphere_aggregation = SphereAggregation( sphere_aggregation = SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,

View file

@ -26,6 +26,9 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
The radius of the sphere in mm. If None, the signal will be extracted The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information (default None). 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 agg_method : str, optional
The aggregation method to use. The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information See :func:`junifer.stats.get_aggfunc_by_name` for more information
@ -53,6 +56,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
self, self,
coords: str, coords: str,
radius: Optional[float] = None, radius: Optional[float] = None,
allow_overlap: bool = False,
agg_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
cor_method: str = "covariance", cor_method: str = "covariance",
@ -62,6 +66,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
) -> None: ) -> None:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.allow_overlap = allow_overlap
if radius is None or radius <= 0: if radius is None or radius <= 0:
raise_error(f"radius should be > 0: provided {radius}") raise_error(f"radius should be > 0: provided {radius}")
super().__init__( super().__init__(
@ -78,6 +83,7 @@ class FunctionalConnectivitySpheres(FunctionalConnectivityBase):
sphere_aggregation = SphereAggregation( sphere_aggregation = SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,

View file

@ -28,6 +28,9 @@ class ReHoSpheres(ReHoBase):
extracted from a single voxel. See extracted from a single voxel. See
:class:`nilearn.maskers.NiftiSpheresMasker` for more information :class:`nilearn.maskers.NiftiSpheresMasker` for more information
(default None). (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 use_afni : bool, optional
Whether to use AFNI for computing. If None, will use AFNI only Whether to use AFNI for computing. If None, will use AFNI only
if available (default None). if available (default None).
@ -93,6 +96,7 @@ class ReHoSpheres(ReHoBase):
self, self,
coords: str, coords: str,
radius: Optional[float] = None, radius: Optional[float] = None,
allow_overlap: bool = False,
use_afni: Optional[bool] = None, use_afni: Optional[bool] = None,
reho_params: Optional[Dict] = None, reho_params: Optional[Dict] = None,
agg_method: str = "mean", agg_method: str = "mean",
@ -102,6 +106,7 @@ class ReHoSpheres(ReHoBase):
) -> None: ) -> None:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.allow_overlap = allow_overlap
self.reho_params = reho_params self.reho_params = reho_params
self.agg_method = agg_method self.agg_method = agg_method
self.agg_method_params = agg_method_params self.agg_method_params = agg_method_params
@ -142,6 +147,7 @@ class ReHoSpheres(ReHoBase):
sphere_aggregation = SphereAggregation( sphere_aggregation = SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,

View file

@ -28,6 +28,9 @@ class SphereAggregation(BaseMarker):
extracted from a single voxel. See extracted from a single voxel. See
:class:`nilearn.maskers.NiftiSpheresMasker` for more information :class:`nilearn.maskers.NiftiSpheresMasker` for more information
(default None). (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 method : str, optional
The aggregation method to use. The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information See :func:`junifer.stats.get_aggfunc_by_name` for more information
@ -54,6 +57,7 @@ class SphereAggregation(BaseMarker):
self, self,
coords: str, coords: str,
radius: Optional[float] = None, radius: Optional[float] = None,
allow_overlap: bool = False,
method: str = "mean", method: str = "mean",
method_params: Optional[Dict[str, Any]] = None, method_params: Optional[Dict[str, Any]] = None,
masks: Union[str, Dict, List[Union[Dict, str]], None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
@ -62,6 +66,7 @@ class SphereAggregation(BaseMarker):
) -> None: ) -> None:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.allow_overlap = allow_overlap
self.method = method self.method = method
self.method_params = method_params or {} self.method_params = method_params or {}
self.masks = masks self.masks = masks
@ -147,6 +152,7 @@ class SphereAggregation(BaseMarker):
masker = JuniferNiftiSpheresMasker( masker = JuniferNiftiSpheresMasker(
seeds=coords, seeds=coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap,
mask_img=mask_img, mask_img=mask_img,
agg_func=agg_func, agg_func=agg_func,
) )

View file

@ -24,6 +24,9 @@ class TemporalSNRSpheres(TemporalSNRBase):
The radius of the sphere in mm. If None, the signal will be extracted The radius of the sphere in mm. If None, the signal will be extracted
from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker` from a single voxel. See :class:`nilearn.maskers.NiftiSpheresMasker`
for more information (default None). 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 agg_method : str, optional
The aggregation method to use. The aggregation method to use.
See :func:`junifer.stats.get_aggfunc_by_name` for more information See :func:`junifer.stats.get_aggfunc_by_name` for more information
@ -45,6 +48,7 @@ class TemporalSNRSpheres(TemporalSNRBase):
self, self,
coords: str, coords: str,
radius: Optional[float] = None, radius: Optional[float] = None,
allow_overlap: bool = False,
agg_method: str = "mean", agg_method: str = "mean",
agg_method_params: Optional[Dict] = None, agg_method_params: Optional[Dict] = None,
masks: Union[str, Dict, List[Union[Dict, str]], None] = None, masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
@ -52,6 +56,7 @@ class TemporalSNRSpheres(TemporalSNRBase):
) -> None: ) -> None:
self.coords = coords self.coords = coords
self.radius = radius self.radius = radius
self.allow_overlap = allow_overlap
if radius is None or radius <= 0: if radius is None or radius <= 0:
raise_error(f"radius should be > 0: provided {radius}") raise_error(f"radius should be > 0: provided {radius}")
super().__init__( super().__init__(
@ -91,6 +96,7 @@ class TemporalSNRSpheres(TemporalSNRBase):
sphere_aggregation = SphereAggregation( sphere_aggregation = SphereAggregation(
coords=self.coords, coords=self.coords,
radius=self.radius, radius=self.radius,
allow_overlap=self.allow_overlap,
method=self.agg_method, method=self.agg_method,
method_params=self.agg_method_params, method_params=self.agg_method_params,
masks=self.masks, masks=self.masks,

View file

@ -28,6 +28,7 @@ def get_aggfunc_by_name(
* ``mean`` -> :func:`numpy.mean` * ``mean`` -> :func:`numpy.mean`
* ``std`` -> :func:`numpy.std` * ``std`` -> :func:`numpy.std`
* ``trim_mean`` -> :func:`scipy.stats.trim_mean` * ``trim_mean`` -> :func:`scipy.stats.trim_mean`
* ``count`` -> :func:`junifer.stats.count`
func_params : dict, optional func_params : dict, optional
Parameters to pass to the function. 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 from functools import partial # local import to avoid sphinx error
# check validity of names # 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: if func_params is None:
func_params = {} func_params = {}
# apply functions # apply functions
@ -74,6 +81,8 @@ def get_aggfunc_by_name(
func = np.std func = np.std
elif name == "trim_mean": elif name == "trim_mean":
func = partial(trim_mean, **func_params) func = partial(trim_mean, **func_params)
elif name == "count":
func = count
else: else:
raise_error( raise_error(
f"Function {name} unknown. Please provide any of " f"Function {name} unknown. Please provide any of "
@ -82,6 +91,24 @@ def get_aggfunc_by_name(
return func return func
synchon commented 2023-03-15 13:21:07 +00:00 (Migrated from github.com)

Should be an int.

Should be an `int`.
def count(data: np.ndarray, axis: int = 0) -> int:
"""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
-------
int
Number of elements along the given axis.
"""
return data.shape[axis]
def winsorized_mean( def winsorized_mean(
data: np.ndarray, axis: Optional[int] = None, **win_params data: np.ndarray, axis: Optional[int] = None, **win_params
) -> np.ndarray: ) -> np.ndarray:

View file

@ -8,7 +8,7 @@ from typing import Dict, Optional
import numpy as np import numpy as np
import pytest import pytest
from junifer.stats import get_aggfunc_by_name, winsorized_mean from junifer.stats import count, get_aggfunc_by_name, winsorized_mean
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -17,6 +17,7 @@ from junifer.stats import get_aggfunc_by_name, winsorized_mean
("winsorized_mean", {"limits": [0.2, 0.7]}), ("winsorized_mean", {"limits": [0.2, 0.7]}),
("mean", None), ("mean", None),
("std", None), ("std", None),
("count", None),
("trim_mean", None), ("trim_mean", None),
("trim_mean", {"proportiontocut": 0.1}), ("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]) input = np.array([22, 4, 9, 8, 5, 3, 7, 2, 1, 6])
output = winsorized_mean(input, limits=[0.1, 0.1]) output = winsorized_mean(input, limits=[0.1, 0.1])
assert output == 5.5 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