Skip to content

Commit eec0aeb

Browse files
authored
Fix dpnp.median/dpnp.nanmedian result shape for empty and size-1 kept dimensions (#3081)
This PR fixes two kept-dimension shape bugs in `dpnp.median`/`dpnp.nanmedian`. **Empty kept dimension.** `dpnp.median` and `dpnp.nanmedian` raised `ValueError` when called with a tuple (or list) `axis` and one of the kept, non-reduced dimensions has size 0, e.g. ``` dpnp.median(dpnp.empty((0, 3, 4)), axis=(1, 2)) ``` The sequence-axis path flattens the reduced axes in `_flatten_array_along_axes` with `a.reshape(kept_shape + (-1,))`. The `-1` cannot be inferred once the array has size 0 and a kept axis is also 0 (every merged length yields a size-0 array, so the dimension is ambiguous), which makes `reshape` raise. The merged length is now computed explicitly with `math.prod(...)`, so the tuple-axis path returns the same empty result as the single-axis path. **Size-1 kept dimension.** `dpnp.nanmedian` dropped kept dimensions of size 1 and returned the wrong shape, e.g. a `(1, 5)` array reduced over `axis=1` returned shape `()` instead of `(1,)`. `_calc_nanmedian` ended with an unconditional `dpnp.squeeze(res)`, which also removed non-reduced size-1 axes; it now squeezes only the reduced trailing axis with `dpnp.squeeze(res, axis=-1)`.
1 parent 6c3fca4 commit eec0aeb

4 files changed

Lines changed: 51 additions & 2 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -110,7 +110,9 @@ This release is compatible with NumPy 2.5.
110110
* Fixed the list of events the copy kernels of `dpnp.reshape`, `dpnp.tensor.reshape`, `dpnp.roll` and `dpnp.tensor.roll` wait on being padded with default-constructed events [#3072](https://github.com/IntelPython/dpnp/pull/3072)
111111
* Fixed `simplify_iteration_three_strides` and `simplify_iteration_four_strides` accumulating into their third and fourth output displacements without zeroing them first, which required the caller to initialize them [#3072](https://github.com/IntelPython/dpnp/pull/3072)
112112
* Fixed `dpnp.ndarray.flat` indexing and assignment edge cases, adding support for slices, ellipsis, and integer/boolean array indices [#3045](https://github.com/IntelPython/dpnp/pull/3045)
113+
* Fixed `dpnp.nanmedian` dropping kept dimensions of size 1, which produced a wrong result shape [#3081](https://github.com/IntelPython/dpnp/pull/3081)
113114
* Fixed incorrect results of `dpnp.tensor.vecdot` in some cases with strided outputs and of `dpnp.tensor` reductions, `dpnp.tensor.vecdot` and `dpnp.tensor.matmul` on large inputs with some data types [#3082](https://github.com/IntelPython/dpnp/pull/3082)
115+
* Fixed `dpnp.median` and `dpnp.nanmedian` raising a `ValueError` for a tuple `axis` when a kept dimension has size 0 [#3081](https://github.com/IntelPython/dpnp/pull/3081)
114116

115117
### Security
116118

‎dpnp/dpnp_utils/dpnp_utils_statistics.py‎

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
# THE POSSIBILITY OF SUCH DAMAGE.
2727
# *****************************************************************************
2828

29+
import math
2930
import warnings
3031

3132
import dpnp
@@ -88,7 +89,8 @@ def _calc_nanmedian(a, out=None):
8889
if mask.all(axis=-1).any():
8990
warnings.warn("All-NaN slice encountered", RuntimeWarning, stacklevel=6)
9091

91-
return dpnp.squeeze(res)
92+
# only drop the reduced axis, keep size-1 dimensions that are not reduced
93+
return dpnp.squeeze(res, axis=-1)
9294

9395

9496
def _flatten_array_along_axes(a, axes_to_flatten, overwrite_input):
@@ -102,7 +104,10 @@ def _flatten_array_along_axes(a, axes_to_flatten, overwrite_input):
102104
# Move the axes_to_flatten to the end
103105
destination = list(range(len(axes_to_keep), a_ndim))
104106
a_moved = dpnp.moveaxis(a, axes_to_flatten, destination)
105-
new_shape = tuple(a.shape[axis] for axis in axes_to_keep) + (-1,)
107+
# Compute the merged length explicitly instead of letting `reshape` infer
108+
# it with -1, since -1 is ambiguous when a kept axis has size 0
109+
merged = math.prod(a.shape[axis] for axis in axes_to_flatten)
110+
new_shape = tuple(a.shape[axis] for axis in axes_to_keep) + (merged,)
106111
a_flatten = a_moved.reshape(new_shape)
107112

108113
# Note that the output of a_flatten is not necessarily a view of the input

‎dpnp/tests/test_nanfunctions.py‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -413,6 +413,34 @@ def test_empty(self, axis, shape):
413413
expected = numpy.nanmedian(a, axis=axis)
414414
assert_dtype_allclose(result, expected)
415415

416+
@pytest.mark.usefixtures("suppress_mean_empty_slice_numpy_warnings")
417+
@pytest.mark.parametrize(
418+
"keepdims, out_shape", [(False, (0,)), (True, (0, 1, 1))]
419+
)
420+
def test_empty_kept_dim(self, keepdims, out_shape):
421+
a = numpy.empty((0, 3, 4))
422+
ia = dpnp.array(a)
423+
424+
result = dpnp.nanmedian(ia, axis=(1, 2), keepdims=keepdims)
425+
assert result.shape == out_shape
426+
if numpy_version() >= "2.5.4":
427+
expected = numpy.nanmedian(a, axis=(1, 2), keepdims=keepdims)
428+
assert_dtype_allclose(result, expected)
429+
430+
@pytest.mark.usefixtures("suppress_mean_empty_slice_numpy_warnings")
431+
@pytest.mark.parametrize(
432+
"shape, axis", [((1, 5), 1), ((3, 1, 4), 2), ((2, 1, 5), (0, 2))]
433+
)
434+
@pytest.mark.parametrize("keepdims", [True, False])
435+
def test_size1_kept_dim(self, shape, axis, keepdims):
436+
a = generate_random_numpy_array(shape)
437+
a.flat[0] = numpy.nan
438+
ia = dpnp.array(a)
439+
440+
result = dpnp.nanmedian(ia, axis=axis, keepdims=keepdims)
441+
expected = numpy.nanmedian(a, axis=axis, keepdims=keepdims)
442+
assert_dtype_allclose(result, expected)
443+
416444
@pytest.mark.parametrize("dtype", get_all_dtypes(no_none=True))
417445
@pytest.mark.parametrize("axis", [None, 0, (-1,), [0, 1], (0, -2, -1)])
418446
def test_no_nan(self, dtype, axis):

‎dpnp/tests/test_statistics.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -944,6 +944,20 @@ def test_empty(self, axis, shape):
944944
expected = numpy.median(a, axis=axis)
945945
assert_dtype_allclose(result, expected)
946946

947+
@pytest.mark.usefixtures("suppress_mean_empty_slice_numpy_warnings")
948+
@pytest.mark.parametrize(
949+
"keepdims, out_shape", [(False, (0,)), (True, (0, 1, 1))]
950+
)
951+
def test_empty_kept_dim(self, keepdims, out_shape):
952+
a = numpy.empty((0, 3, 4))
953+
ia = dpnp.array(a)
954+
955+
result = dpnp.median(ia, axis=(1, 2), keepdims=keepdims)
956+
assert result.shape == out_shape
957+
if numpy_version() >= "2.5.4":
958+
expected = numpy.median(a, axis=(1, 2), keepdims=keepdims)
959+
assert_dtype_allclose(result, expected)
960+
947961
@pytest.mark.parametrize("dtype", get_all_dtypes())
948962
@pytest.mark.parametrize(
949963
"axis, out_shape", [(0, (3,)), (1, (2,)), ((0, 1), ())]

0 commit comments

Comments
 (0)