Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/api-assorted.md
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
kron
nan_to_num
nanmax
nanmean
nanmin
nansum
nunique
Expand Down
2 changes: 2 additions & 0 deletions src/array_api_extra/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
kron,
nan_to_num,
nanmax,
nanmean,
nanmin,
nansum,
nunique,
Expand Down Expand Up @@ -61,6 +62,7 @@
"lazy_apply",
"nan_to_num",
"nanmax",
"nanmean",
"nanmin",
"nansum",
"nunique",
Expand Down
53 changes: 53 additions & 0 deletions src/array_api_extra/_delegation.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
"nan_to_num",
"nanmax",
"nanmin",
"nanmean",
"nansum",
"nunique",
"one_hot",
Expand Down Expand Up @@ -1878,3 +1879,55 @@ def nansum(
return xp.nansum(a, axis=axis)

return _funcs.nansum(a, axis=axis, xp=xp)


def nanmean(
a: Array,
/,
*,
axis: int | tuple[int, ...] | None = None,
xp: ArrayNamespace | None = None,
) -> Array:
"""
Return the mean of the array elements along a given axis, ignoring NaNs.

Parameters
----------
a : Array
Input array.
axis : int or tuple of ints or None, optional
Axis or axes along which the mean is computed. The default is to compute
the mean of the flattened array.
xp : array_namespace, optional
The standard-compatible namespace for `a`. Default: infer.

Returns
-------
array
An array of mean values along the given axis, ignoring NaNs.

Examples
--------
>>> import array_api_extra as xpx
>>> import array_api_strict as xp
>>> a = xp.asarray([[5, 3, xp.nan, 1], [4, xp.nan, 2, xp.nan]])
>>> xpx.nanmean(a)
Array(3., dtype=array_api_strict.float64)
>>> xpx.nanmean(a, axis=0)
Array([4.5, 3., 2.5, 1.], dtype=array_api_strict.float64)
>>> xpx.nanmean(a, axis=1)
Array([3., 3.], dtype=array_api_strict.float64)
"""
if xp is None:
xp = array_namespace(a)

if (
is_numpy_namespace(xp)
or is_cupy_namespace(xp)
or is_dask_namespace(xp)
or is_jax_namespace(xp)
or is_torch_namespace(xp)
):
return xp.nanmean(a, axis=axis)

return _funcs.nanmean(a, axis=axis, xp=xp)
29 changes: 29 additions & 0 deletions src/array_api_extra/_lib/_funcs.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
"kron",
"nan_to_num",
"nanmax",
"nanmean",
"nanmin",
"nansum",
"nunique",
Expand Down Expand Up @@ -871,3 +872,31 @@ def nansum( # numpydoc ignore=PR01,RT01
device_a = _compat.device(a)
zero = xp.asarray(0, dtype=a.dtype, device=device_a)
return xp.sum(xp.where(mask, zero, a), axis=axis)


def nanmean( # numpydoc ignore=PR01,RT01
a: Array,
/,
*,
axis: int | tuple[int, ...] | None,
xp: ArrayNamespace,
) -> Array:
"""See docstring in `array_api_extra._delegation.py`."""
mask = xp.isnan(a)
device_a = _compat.device(a)
zero = xp.asarray(0, dtype=a.dtype, device=device_a)
sum_ = xp.sum(xp.where(mask, zero, a), axis=axis)
count = xp.count_nonzero(~mask, axis=axis)
safe_count = xp.astype(
xp.where(count == 0, xp.asarray(1, dtype=count.dtype, device=device_a), count),
sum_.dtype,
copy=False,
)
result = sum_ / safe_count
if xp.any(count == 0):
result = xp.where(
count == 0,
xp.asarray(xp.nan, dtype=result.dtype, device=device_a),
result,
)
return result
65 changes: 65 additions & 0 deletions tests/test_funcs.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
kron,
nan_to_num,
nanmax,
nanmean,
nanmin,
nansum,
nunique,
Expand Down Expand Up @@ -75,6 +76,7 @@
lazy_xp_function(isin)
lazy_xp_function(kron)
lazy_xp_function(nan_to_num)
lazy_xp_function(nanmean)
lazy_xp_function(nansum)
lazy_xp_function(nunique)
lazy_xp_function(one_hot)
Expand Down Expand Up @@ -2531,3 +2533,66 @@ def test_xp(self, axis: int | None, expected_list: list[float], xp: ArrayNamespa
res = nansum(a, axis=axis, xp=xp)
expected = xp.asarray(expected_list)
assert_equal(res, expected)


class TestNanMean:
def test_simple(self, xp: ArrayNamespace):
a = xp.asarray([[1.0, 2.0], [3.0, xp.nan]])

res = nanmean(a)
assert res == 2.0

res = nanmean(a, axis=0)
expected = xp.asarray([2.0, 2.0])
assert_equal(res, expected)

res = nanmean(a, axis=1)
expected = xp.asarray([1.5, 3.0])
assert_equal(res, expected)

def test_bigger(self, xp: ArrayNamespace):
a = xp.asarray(
[
[1.0, xp.nan, 4.0, 5.0],
[xp.nan, -2.0, xp.nan, -4.0],
[2.0, 1.0, 3.0, xp.nan],
]
)

res = nanmean(a, axis=0)
expected = xp.asarray([1.5, -0.5, 3.5, 0.5])
assert_equal(res, expected)

res = nanmean(a, axis=1)
expected = xp.asarray([3.3333333, -3.0, 2.0])
assert_close(res, expected)

@pytest.mark.filterwarnings("ignore:.*Mean of empty slice.*:RuntimeWarning")
def test_all_nan_slice(self, xp: ArrayNamespace):
a = xp.asarray([[xp.nan, 1.0], [xp.nan, xp.nan]])

res = nanmean(a, axis=0, xp=xp)
expected = xp.asarray([xp.nan, 1.0])
assert_equal(res, expected)

def test_scalar(self, xp: ArrayNamespace):
a = xp.asarray(1.0)
assert nanmean(a) == 1.0

@pytest.mark.skip_xp_backend(
Backend.TORCH, reason="torch.nanmean does not support tensors on meta device"
)
@pytest.mark.parametrize("axis", [None, 0, 1])
def test_device(self, axis: int | None, xp: ArrayNamespace, device: Device):
a = xp.asarray([[4.0, xp.nan, 1.0], [2.0, 5.0, xp.nan]], device=device)
res = nanmean(a, axis=axis)
assert get_device(res) == device

@pytest.mark.parametrize(
("axis", "expected_list"), [(0, [3.0, 5.0, 1.0]), (1, [2.5, 3.5])]
)
def test_xp(self, axis: int | None, expected_list: list[float], xp: ArrayNamespace):
a = xp.asarray([[4.0, xp.nan, 1.0], [2.0, 5.0, xp.nan]])
res = nanmean(a, axis=axis, xp=xp)
expected = xp.asarray(expected_list)
assert_equal(res, expected)
Loading