[ENH]: Add mode as an aggregation function for get_aggfunc_by_name #287
5 changed files with 25 additions and 4 deletions
1
docs/changes/newsfragments/287.enh
Normal file
1
docs/changes/newsfragments/287.enh
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Add ``mode`` as an aggregation function option in :func:`.get_aggfunc_by_name` by `Synchon Mandal`_
|
||||||
|
|
@ -35,6 +35,7 @@ def test_get_dependency_information_short() -> None:
|
||||||
assert list(dependency_information.keys()) == [
|
assert list(dependency_information.keys()) == [
|
||||||
"click",
|
"click",
|
||||||
"numpy",
|
"numpy",
|
||||||
|
"scipy",
|
||||||
"datalad",
|
"datalad",
|
||||||
"pandas",
|
"pandas",
|
||||||
"nibabel",
|
"nibabel",
|
||||||
|
|
@ -51,6 +52,7 @@ def test_get_dependency_information_long() -> None:
|
||||||
for key in [
|
for key in [
|
||||||
"click",
|
"click",
|
||||||
"numpy",
|
"numpy",
|
||||||
|
"scipy",
|
||||||
"datalad",
|
"datalad",
|
||||||
"pandas",
|
"pandas",
|
||||||
"nibabel",
|
"nibabel",
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,7 @@
|
||||||
from typing import Any, Callable, Dict, List, Optional
|
from typing import Any, Callable, Dict, List, Optional
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from scipy.stats import trim_mean
|
from scipy.stats import mode, trim_mean
|
||||||
from scipy.stats.mstats import winsorize
|
from scipy.stats.mstats import winsorize
|
||||||
|
|
||||||
from .utils import logger, raise_error
|
from .utils import logger, raise_error
|
||||||
|
|
@ -24,10 +24,11 @@ def get_aggfunc_by_name(
|
||||||
Name to identify the function. Currently supported names and
|
Name to identify the function. Currently supported names and
|
||||||
corresponding functions are:
|
corresponding functions are:
|
||||||
|
|
||||||
* ``winsorized_mean`` -> :func:`scipy.stats.mstats.winsorize`
|
|
||||||
* ``mean`` -> :func:`numpy.mean`
|
* ``mean`` -> :func:`numpy.mean`
|
||||||
* ``std`` -> :func:`numpy.std`
|
* ``winsorized_mean`` -> :func:`scipy.stats.mstats.winsorize`
|
||||||
* ``trim_mean`` -> :func:`scipy.stats.trim_mean`
|
* ``trim_mean`` -> :func:`scipy.stats.trim_mean`
|
||||||
|
* ``mode`` -> :func:`scipy.stats.mode`
|
||||||
|
* ``std`` -> :func:`numpy.std`
|
||||||
* ``count`` -> :func:`.count`
|
* ``count`` -> :func:`.count`
|
||||||
* ``select`` -> :func:`.select`
|
* ``select`` -> :func:`.select`
|
||||||
|
|
||||||
|
|
@ -40,6 +41,7 @@ def get_aggfunc_by_name(
|
||||||
-------
|
-------
|
||||||
function
|
function
|
||||||
Respective function with ``func_params`` parameter set.
|
Respective function with ``func_params`` parameter set.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
from functools import partial # local import to avoid sphinx error
|
from functools import partial # local import to avoid sphinx error
|
||||||
|
|
||||||
|
|
@ -51,6 +53,7 @@ def get_aggfunc_by_name(
|
||||||
"trim_mean",
|
"trim_mean",
|
||||||
"count",
|
"count",
|
||||||
"select",
|
"select",
|
||||||
|
"mode",
|
||||||
}
|
}
|
||||||
if func_params is None:
|
if func_params is None:
|
||||||
func_params = {}
|
func_params = {}
|
||||||
|
|
@ -93,6 +96,8 @@ def get_aggfunc_by_name(
|
||||||
elif pick is not None and drop is not None:
|
elif pick is not None and drop is not None:
|
||||||
raise_error("Either pick or drop must be specified, not both.")
|
raise_error("Either pick or drop must be specified, not both.")
|
||||||
func = partial(select, **func_params)
|
func = partial(select, **func_params)
|
||||||
|
elif name == "mode":
|
||||||
|
func = partial(mode, **func_params)
|
||||||
else:
|
else:
|
||||||
raise_error(
|
raise_error(
|
||||||
f"Function {name} unknown. Please provide any of "
|
f"Function {name} unknown. Please provide any of "
|
||||||
|
|
@ -115,6 +120,7 @@ def count(data: np.ndarray, axis: int = 0) -> np.ndarray:
|
||||||
-------
|
-------
|
||||||
numpy.ndarray
|
numpy.ndarray
|
||||||
Number of elements along the given axis.
|
Number of elements along the given axis.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
ax_size = data.shape[axis]
|
ax_size = data.shape[axis]
|
||||||
if axis < 0:
|
if axis < 0:
|
||||||
|
|
@ -137,7 +143,7 @@ def winsorized_mean(
|
||||||
The axis to calculate winsorized mean on (default None).
|
The axis to calculate winsorized mean on (default None).
|
||||||
**win_params : dict
|
**win_params : dict
|
||||||
Dictionary containing the keyword arguments for the winsorize function.
|
Dictionary containing the keyword arguments for the winsorize function.
|
||||||
E.g. ``{'limits': [0.1, 0.1]}``.
|
E.g., ``{'limits': [0.1, 0.1]}``.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
|
|
@ -149,6 +155,7 @@ def winsorized_mean(
|
||||||
--------
|
--------
|
||||||
scipy.stats.mstats.winsorize :
|
scipy.stats.mstats.winsorize :
|
||||||
The winsorize function used in this function.
|
The winsorize function used in this function.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
win_dat = winsorize(data, axis=axis, **win_params)
|
win_dat = winsorize(data, axis=axis, **win_params)
|
||||||
win_mean = win_dat.mean(axis=axis)
|
win_mean = win_dat.mean(axis=axis)
|
||||||
|
|
@ -180,6 +187,13 @@ def select(
|
||||||
numpy.ndarray
|
numpy.ndarray
|
||||||
Subset of the inputted data with the select settings
|
Subset of the inputted data with the select settings
|
||||||
applied as specified in ``select_params``.
|
applied as specified in ``select_params``.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
ValueError
|
||||||
|
If both ``pick`` and ``drop`` are None or
|
||||||
|
if both ``pick`` and ``drop`` are not None.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if pick is None and drop is None:
|
if pick is None and drop is None:
|
||||||
|
|
|
||||||
|
|
@ -21,6 +21,8 @@ from junifer.stats import count, get_aggfunc_by_name, select, winsorized_mean
|
||||||
("count", None),
|
("count", None),
|
||||||
("trim_mean", None),
|
("trim_mean", None),
|
||||||
("trim_mean", {"proportiontocut": 0.1}),
|
("trim_mean", {"proportiontocut": 0.1}),
|
||||||
|
("mode", None),
|
||||||
|
("mode", {"keepdims": True}),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_get_aggfunc_by_name(name: str, params: Optional[Dict]) -> None:
|
def test_get_aggfunc_by_name(name: str, params: Optional[Dict]) -> None:
|
||||||
|
|
|
||||||
|
|
@ -37,6 +37,7 @@ classifiers = [
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"click>=8.1.3,<8.2",
|
"click>=8.1.3,<8.2",
|
||||||
"numpy>=1.24,<1.27",
|
"numpy>=1.24,<1.27",
|
||||||
|
"scipy>=1.9.0,<=1.11.4",
|
||||||
"datalad>=0.15.4,<0.20",
|
"datalad>=0.15.4,<0.20",
|
||||||
"pandas>=1.4.0,<2.2",
|
"pandas>=1.4.0,<2.2",
|
||||||
"nibabel>=3.2.0,<5.11",
|
"nibabel>=3.2.0,<5.11",
|
||||||
|
|
@ -177,6 +178,7 @@ known-first-party = ["junifer"]
|
||||||
known-third-party =[
|
known-third-party =[
|
||||||
"click",
|
"click",
|
||||||
"numpy",
|
"numpy",
|
||||||
|
"scipy",
|
||||||
"datalad",
|
"datalad",
|
||||||
"pandas",
|
"pandas",
|
||||||
"nibabel",
|
"nibabel",
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue