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 CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,7 @@ This release is compatible with NumPy 2.5.
* Fixed `dpnp.all` and `dpnp.any` aborting when reducing over an empty axis (e.g. an array with a zero-length dimension) [#3021](https://github.com/IntelPython/dpnp/pull/3021)
* Released the GIL before the blocking OneMKL DFT calls in the FFT extension [#3040](https://github.com/IntelPython/dpnp/pull/3040)
* Fixed `astype` casting an out-of-range floating point value to a signed narrow integer type saturating to the destination min/max instead of wrapping like NumPy, generalizing the earlier unsigned-only fix [#3033](https://github.com/IntelPython/dpnp/pull/3033)
* Fixed `dpnp.ndarray.flat` indexing edge cases, adding support for slices, ellipsis, and integer/boolean array indices [#3045](https://github.com/IntelPython/dpnp/pull/3045)

### Security

Expand Down
34 changes: 31 additions & 3 deletions dpnp/dpnp_array.py
Original file line number Diff line number Diff line change
Expand Up @@ -1332,9 +1332,37 @@ def flags(self):
@property
def flat(self):
"""
Return a flat iterator, or set a flattened version of self to value.
A 1-D iterator over the array.

""" # noqa: D200
This is a :obj:`dpnp.flatiter` instance, which acts similarly to, but
is not a subclass of, Python's built-in iterator object.

For full documentation refer to :obj:`numpy.ndarray.flat`.

See Also
--------
:obj:`dpnp.flatiter` : Flat iterator object to iterate over arrays.
:obj:`dpnp.ndarray.flatten` : Return a flattened copy of the array.

Examples
--------
>>> import dpnp as np
>>> x = np.arange(1, 7).reshape(2, 3)
>>> x
array([[1, 2, 3],
[4, 5, 6]])
>>> x.flat[3]
array(4)
>>> x.T.flat[3]
array(5)

An assignment example:

>>> x.flat[[1, 4]] = 1; x
array([[1, 1, 3],
[4, 1, 6]])

"""

return dpnp.flatiter(self)

Expand Down Expand Up @@ -1367,7 +1395,7 @@ def flatten(self, /, order="C"):
See Also
--------
:obj:`dpnp.ravel` : Return a flattened array.
:obj:`dpnp.flat` : A 1-D flat iterator over the array.
:obj:`dpnp.ndarray.flat` : A 1-D flat iterator over the array.

Examples
--------
Expand Down
162 changes: 121 additions & 41 deletions dpnp/dpnp_flatiter.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,59 +32,139 @@


class flatiter:
"""Flat iterator object to iterate over arrays."""
"""
Flat iterator object to iterate over arrays.

def __init__(self, X):
if type(X) is not dpnp.ndarray:
A flat iterator is returned by :obj:`dpnp.ndarray.flat` for any array. It
allows iterating over the array as if it were a 1-D array, either in a
for-loop or by calling its ``next`` method.

Iteration is done in row-major, C-style order (the last index varying the
fastest). The iterator can also be indexed using basic slicing or advanced
indexing.

For full documentation refer to :obj:`numpy.flatiter`.

See Also
--------
:obj:`dpnp.ndarray.flat` : Return a flat iterator over an array.
:obj:`dpnp.ndarray.flatten` : Return a flattened copy of an array.

Examples
--------
>>> import dpnp as np
>>> x = np.arange(6).reshape(2, 3)
>>> for item in x.flat:
... print(item)
0
1
2
3
4
5

>>> x.flat[2:4]
array([2, 3])

"""

def __init__(self, a):
if not isinstance(a, dpnp.ndarray):
raise TypeError(
"Argument must be of type dpnp.ndarray, got {}".format(type(X))
f"An array must be of type dpnp.ndarray, but got {type(a)}"
)
self._arr = a
self._size = a.size
self._i = 0

@staticmethod
def _reject_newaxis(key):
# newaxis (None) is valid for array indexing but not for flat indexing
if key is None or (
isinstance(key, tuple) and any(k is None for k in key)
):
raise IndexError(
"only integers, slices (`:`), ellipsis (`...`) and integer "
"or boolean arrays are valid indices"
)
self.arr_ = X
self.size_ = X.size
self.i_ = 0

def _multiindex(self, i):
nd = self.arr_.ndim
if nd == 0:
if i == 0:
return ()
raise KeyError
elif nd == 1:
return (i,)
sh = self.arr_.shape
i_ = i
multi_index = [0] * nd
for k in reversed(range(1, nd)):
si = sh[k]
q = i_ // si
multi_index[k] = i_ - q * si
i_ = q
multi_index[0] = i_
return tuple(multi_index)

def _check_bounds(self, key):
# fancy int indices wrap instead of raising, so check them vs NumPy
if key is Ellipsis or isinstance(key, (slice, bool, tuple)):
return

if isinstance(key, int) or (
callable(getattr(key, "__index__", None))
and not hasattr(key, "ndim")
):
return # scalar int: regular indexing checks it

try:
idx = dpnp.asarray(key, sycl_queue=self._arr.sycl_queue)
except Exception:
return # let regular indexing raise

if idx.dtype.kind not in "iu" or idx.size == 0:
return

size = self._size
hi, lo = int(dpnp.max(idx)), int(dpnp.min(idx))
if hi >= size:
raise IndexError(f"index {hi} is out of bounds for size {size}")
if lo < -size:
raise IndexError(f"index {lo} is out of bounds for size {size}")

def _flatten(self):
# C-order flat view (copy if non-contiguous)
return dpnp.reshape(self._arr, -1)

def __getitem__(self, key):
idx = getattr(key, "__index__", None)
if not callable(idx):
raise TypeError(key)
i = idx()
mi = self._multiindex(i)
return self.arr_.__getitem__(mi)
self._reject_newaxis(key)
self._check_bounds(key)

# flat always yields a copy, never a view
return self._flatten()[key].copy()

def __setitem__(self, key, val):
idx = getattr(key, "__index__", None)
if not callable(idx):
raise TypeError(key)
i = idx()
mi = self._multiindex(i)
return self.arr_.__setitem__(mi, val)
self._reject_newaxis(key)
self._check_bounds(key)

if isinstance(key, tuple) and len(key) == 0:
# NumPy rejects arr.flat[()] = val
raise IndexError(
"Assigning to a flat iterator with a 0-D index is not "
"supported"
)

a = self._arr
exec_q = a.sycl_queue
usm_type = a.usm_type

# resolve key to flat positions, reusing regular indexing to validate
flat_index = dpnp.arange(a.size, sycl_queue=exec_q, usm_type=usm_type)
idx = flat_index[key]

if not dpnp.isscalar(val):
val = dpnp.asarray(
val, sycl_queue=exec_q, usm_type=usm_type
).ravel()
n = idx.size
if 0 < val.size != n:
# cycles the values over the selection
val = val[
dpnp.arange(n, sycl_queue=exec_q, usm_type=usm_type)
% val.size
]

dpnp.put(a, idx, val)

def __iter__(self):
return self

def __next__(self):
if self.i_ < self.size_:
val = self.__getitem__(self.i_)
self.i_ = self.i_ + 1
if self._i < self._size:
val = self.__getitem__(self._i)
self._i = self._i + 1
return val
else:
raise StopIteration
Loading
Loading