From 937f1a82318472a6f74a40d07e7570dc483069e8 Mon Sep 17 00:00:00 2001 From: Ushnah Abbasi <90151235+UshnahAbbasi@users.noreply.github.com> Date: Wed, 12 Aug 2026 17:36:55 +0500 Subject: [PATCH] implemented nanmean --- docs/api-assorted.md | 1 + src/array_api_extra/__init__.py | 2 + src/array_api_extra/_delegation.py | 53 ++++++++++++++++++++++++ src/array_api_extra/_lib/_funcs.py | 29 +++++++++++++ tests/test_funcs.py | 65 ++++++++++++++++++++++++++++++ 5 files changed, 150 insertions(+) diff --git a/docs/api-assorted.md b/docs/api-assorted.md index 26ccb4d7..fb390c09 100644 --- a/docs/api-assorted.md +++ b/docs/api-assorted.md @@ -23,6 +23,7 @@ kron nan_to_num nanmax + nanmean nanmin nansum nunique diff --git a/src/array_api_extra/__init__.py b/src/array_api_extra/__init__.py index e0a5a210..e597a611 100644 --- a/src/array_api_extra/__init__.py +++ b/src/array_api_extra/__init__.py @@ -15,6 +15,7 @@ kron, nan_to_num, nanmax, + nanmean, nanmin, nansum, nunique, @@ -61,6 +62,7 @@ "lazy_apply", "nan_to_num", "nanmax", + "nanmean", "nanmin", "nansum", "nunique", diff --git a/src/array_api_extra/_delegation.py b/src/array_api_extra/_delegation.py index e446d35d..a176d0a9 100644 --- a/src/array_api_extra/_delegation.py +++ b/src/array_api_extra/_delegation.py @@ -40,6 +40,7 @@ "nan_to_num", "nanmax", "nanmin", + "nanmean", "nansum", "nunique", "one_hot", @@ -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) diff --git a/src/array_api_extra/_lib/_funcs.py b/src/array_api_extra/_lib/_funcs.py index 2d89bf44..e7da36af 100644 --- a/src/array_api_extra/_lib/_funcs.py +++ b/src/array_api_extra/_lib/_funcs.py @@ -39,6 +39,7 @@ "kron", "nan_to_num", "nanmax", + "nanmean", "nanmin", "nansum", "nunique", @@ -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 diff --git a/tests/test_funcs.py b/tests/test_funcs.py index f5d756b0..4149708d 100644 --- a/tests/test_funcs.py +++ b/tests/test_funcs.py @@ -31,6 +31,7 @@ kron, nan_to_num, nanmax, + nanmean, nanmin, nansum, nunique, @@ -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) @@ -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)