使用ndenumerate遍历带切片的二维numpy数组时索引错误的原因
为什么用np.ndenumerate遍历切片后的二维numpy数组会得到“错误”索引?
先看你的代码和输出情况:
代码示例:
import numpy as np arr = np.array([[1, 2, 3, 4], [5, 6, 7, 8]]) for idx, x in np.ndenumerate(arr[:, ::2]): print(idx, x)
实际输出:
(0, 0) 1 (0, 1) 3 (1, 0) 5 (1, 1) 7
预期输出:
(0, 0) 1 (0, 2) 3 (1, 0) 5 (1, 2) 7
原因解析
这不是np.ndenumerate的bug,而是它的核心设计逻辑:它只会基于传入的当前数组的形状生成索引,完全不关心这个数组是不是原数组的切片/视图。
你执行arr[:, ::2]后,得到的是一个形状为(2,2)的新数组(虽然它是原数组的视图、共享数据,但形状已经改变)。np.ndenumerate遍历的是这个(2,2)的数组,所以输出的索引是针对这个新数组的位置,而非原数组的位置。
当你不使用切片时,遍历的是原数组本身(形状(2,4)),所以索引自然对应原数组的真实位置。
解决方案:获取原数组的索引
如果你需要拿到元素在原数组中的索引,有几种可行的方法:
方法一:手动映射切片索引
先提取切片对应的原数组列索引,再结合行索引遍历:
import numpy as np arr = np.array([[1, 2, 3, 4], [5, 6, 7, 8]]) # 获取切片对应的原数组列索引 col_indices = np.arange(arr.shape[1])[::2] for row_idx in range(arr.shape[0]): for _, orig_col_idx in enumerate(col_indices): print((row_idx, orig_col_idx), arr[row_idx, orig_col_idx])
方法二:生成原数组的索引矩阵
利用np.indices生成原数组的索引矩阵,再通过切片拿到对应位置的原索引,最后遍历:
import numpy as np arr = np.array([[1, 2, 3, 4], [5, 6, 7, 8]]) slice_obj = (slice(None), slice(None, None, 2)) # 定义切片对象 # 获取切片对应的原数组索引 orig_indices = np.indices(arr.shape)[:, slice_obj] # 扁平化后遍历 for idx, val in zip(zip(orig_indices[0].ravel(), orig_indices[1].ravel()), arr[slice_obj].ravel()): print(idx, val)
方法三:用np.nditer(更灵活)
通过np.nditer的参数直接关联原数组索引:
import numpy as np arr = np.array([[1, 2, 3, 4], [5, 6, 7, 8]]) slice_arr = arr[:, ::2] for it in np.nditer(slice_arr, flags=['multi_index']): # 拿到切片数组的索引后,映射回原数组的列索引 row_idx, slice_col_idx = it.multi_index orig_col_idx = slice_col_idx * 2 # 对应切片步长为2的规则 print((row_idx, orig_col_idx), it[()])
以上方法都能输出你预期的原数组索引。
内容的提问来源于stack exchange,提问作者PatilUdayV
相关产品推荐
相关产品推荐

