|
30 | 30 | isin, |
31 | 31 | kron, |
32 | 32 | nan_to_num, |
| 33 | + nanmin, |
33 | 34 | nunique, |
34 | 35 | one_hot, |
35 | 36 | pad, |
@@ -2203,3 +2204,84 @@ def test_xp(self, xp: ModuleType): |
2203 | 2204 | res = unravel_index(indices, shape, xp=xp) |
2204 | 2205 | for res_arr, exp_arr in zip(res, expected, strict=True): |
2205 | 2206 | 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