Skip to content

Commit dcd9242

Browse files
authored
ENH: add nanmin (#804)
1 parent cf3431d commit dcd9242

5 files changed

Lines changed: 159 additions & 0 deletions

File tree

docs/api-assorted.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
isin
2222
kron
2323
nan_to_num
24+
nanmin
2425
nunique
2526
one_hot
2627
pad

src/array_api_extra/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
isin,
1414
kron,
1515
nan_to_num,
16+
nanmin,
1617
one_hot,
1718
pad,
1819
partition,
@@ -54,6 +55,7 @@
5455
"kron",
5556
"lazy_apply",
5657
"nan_to_num",
58+
"nanmin",
5759
"nunique",
5860
"one_hot",
5961
"pad",

src/array_api_extra/_delegation.py

Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@
3535
"isin",
3636
"kron",
3737
"nan_to_num",
38+
"nanmin",
3839
"one_hot",
3940
"pad",
4041
"partition",
@@ -1583,3 +1584,54 @@ def unravel_index(
15831584
return xp.unravel_index(indices, shape)
15841585

15851586
return _funcs.unravel_index(indices, shape)
1587+
1588+
1589+
def nanmin(
1590+
a: Array,
1591+
/,
1592+
*,
1593+
axis: int | tuple[int, ...] | None = None,
1594+
xp: ModuleType | None = None,
1595+
) -> Array:
1596+
"""
1597+
Return the minimum of the array elements along a given axis, ignoring NaNs.
1598+
1599+
Parameters
1600+
----------
1601+
a : Array
1602+
Input array.
1603+
axis : int or tuple of ints or None, optional
1604+
Axis or axes along which the minimum is computed. The default is to compute
1605+
the minimum of the flattened array.
1606+
xp : array_namespace, optional
1607+
The standard-compatible namespace for `a`. Default: infer.
1608+
1609+
Returns
1610+
-------
1611+
array
1612+
An array of minimum values along the given axis, ignoring NaNs.
1613+
1614+
Examples
1615+
--------
1616+
>>> import array_api_extra as xpx
1617+
>>> import array_api_strict as xp
1618+
>>> a = xp.asarray([[5, 3, xp.nan, 1], [4, xp.nan, 2, xp.nan]])
1619+
>>> xpx.nanmin(a)
1620+
Array(1., dtype=array_api_strict.float64)
1621+
>>> xpx.nanmin(a, axis=0)
1622+
Array([4., 3., 2., 1.], dtype=array_api_strict.float64)
1623+
>>> xpx.nanmin(a, axis=1)
1624+
Array([1., 2.], dtype=array_api_strict.float64)
1625+
"""
1626+
if xp is None:
1627+
xp = array_namespace(a)
1628+
1629+
if (
1630+
is_numpy_namespace(xp)
1631+
or is_cupy_namespace(xp)
1632+
or is_dask_namespace(xp)
1633+
or is_jax_namespace(xp)
1634+
):
1635+
return xp.nanmin(a, axis=axis)
1636+
1637+
return _funcs.nanmin(a, axis=axis, xp=xp)

src/array_api_extra/_lib/_funcs.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@
3838
"isin",
3939
"kron",
4040
"nan_to_num",
41+
"nanmin",
4142
"nunique",
4243
"one_hot",
4344
"pad",
@@ -823,3 +824,24 @@ def unravel_index(indices: Array, shape: tuple[int, ...], /) -> tuple[Array, ...
823824
coords.append(indices % dim)
824825
indices = indices // dim
825826
return tuple(reversed(coords))
827+
828+
829+
def nanmin( # numpydoc ignore=PR01,RT01
830+
a: Array,
831+
/,
832+
*,
833+
axis: int | tuple[int, ...] | None,
834+
xp: ModuleType,
835+
) -> Array:
836+
"""See docstring in `array_api_extra._delegation.py`."""
837+
mask = xp.isnan(a)
838+
device_a = _compat.device(a)
839+
x = xp.min(
840+
xp.where(mask, xp.asarray(+xp.inf, dtype=a.dtype, device=device_a), a),
841+
axis=axis,
842+
)
843+
# Replace Infs from all NaN slices with NaN again
844+
mask = xp.all(mask, axis=axis)
845+
if xp.any(mask):
846+
x = xp.where(mask, xp.asarray(xp.nan, dtype=x.dtype, device=device_a), x)
847+
return x

tests/test_funcs.py

Lines changed: 82 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
isin,
3131
kron,
3232
nan_to_num,
33+
nanmin,
3334
nunique,
3435
one_hot,
3536
pad,
@@ -2203,3 +2204,84 @@ def test_xp(self, xp: ModuleType):
22032204
res = unravel_index(indices, shape, xp=xp)
22042205
for res_arr, exp_arr in zip(res, expected, strict=True):
22052206
assert_equal(res_arr, exp_arr)
2207+
2208+
2209+
class TestNanMin:
2210+
def test_simple(self, xp: ModuleType):
2211+
a = xp.asarray([[1, 2], [3, xp.nan]])
2212+
2213+
# with the default `axis=None` a single scalar is returned
2214+
res = nanmin(a)
2215+
expected = 1.0
2216+
assert res == expected
2217+
2218+
res = nanmin(a, axis=0)
2219+
expected = xp.asarray([1.0, 2.0])
2220+
assert_equal(res, expected)
2221+
2222+
res = nanmin(a, axis=1)
2223+
expected = xp.asarray([1.0, 3.0])
2224+
assert_equal(res, expected)
2225+
2226+
def test_bigger(self, xp: ModuleType):
2227+
a = xp.asarray(
2228+
[
2229+
[1, xp.nan, 4, 5],
2230+
[xp.nan, -2, xp.nan, -4],
2231+
[2, 1, 3, xp.nan],
2232+
]
2233+
)
2234+
2235+
res = nanmin(a, axis=0)
2236+
expected = xp.asarray([1.0, -2.0, 3.0, -4.0])
2237+
assert_equal(res, expected)
2238+
2239+
res = nanmin(a, axis=1)
2240+
expected = xp.asarray([1.0, -4.0, 1.0])
2241+
assert_equal(res, expected)
2242+
2243+
def test_with_infinity(self, xp: ModuleType):
2244+
a = xp.asarray([0.1, 1.0, xp.nan, xp.inf])
2245+
res = nanmin(a)
2246+
expected = 0.1
2247+
assert res == expected
2248+
2249+
a = xp.asarray([0.1, 1.0, xp.nan, -xp.inf])
2250+
res = nanmin(a)
2251+
expected = -xp.inf
2252+
assert res == expected
2253+
2254+
def test_scalar(self, xp: ModuleType):
2255+
a = xp.asarray(1.0)
2256+
assert nanmin(a) == 1.0
2257+
2258+
@pytest.mark.filterwarnings("ignore:.*All-NaN slice*.:RuntimeWarning")
2259+
def test_all_nan_slice_2d(self, xp: ModuleType):
2260+
a = xp.asarray(
2261+
[
2262+
[xp.nan, 5.0],
2263+
[xp.nan, 2.0],
2264+
]
2265+
)
2266+
2267+
res = nanmin(a, axis=0, xp=xp)
2268+
expected = xp.asarray([xp.nan, 2.0])
2269+
assert_equal(res, expected)
2270+
2271+
@pytest.mark.skip_xp_backend(
2272+
Backend.TORCH, reason="torch.nanmin does not support tensors on meta device"
2273+
)
2274+
@pytest.mark.parametrize("axis", [None, 0, 1])
2275+
def test_device(self, axis: int | None, xp: ModuleType, device: Device):
2276+
a = xp.asarray([[4, xp.nan, 1], [2, 5, xp.nan]], device=device)
2277+
res = nanmin(a, axis=axis)
2278+
assert get_device(res) == device
2279+
2280+
@pytest.mark.parametrize(
2281+
("axis", "expected_list"), [(0, [2.0, 3.0, 1.0]), (1, [1.0, 2.0])]
2282+
)
2283+
def test_xp(self, axis: int | None, expected_list: list[float], xp: ModuleType):
2284+
a = xp.asarray([[4, xp.nan, 1], [2, 3, xp.nan]])
2285+
res = nanmin(a, axis=axis, xp=xp)
2286+
expected = xp.asarray(expected_list)
2287+
assert_equal(res, expected)

0 commit comments

Comments
 (0)