NumPy数组切片如何同时返回切片结果与对应原数组索引
NumPy切片同时获取结果数组与对应原数组索引
实现代码
通过索引网格匹配相同切片规则的方式即可实现需求,代码完全匹配给出的校验逻辑与预期输出:
import numpy as np original = np.array([ [5, 3, 7, 3, 2], [8, 4, 22, 6, 4], ]) sliced_array = original[:,::3] # 生成原数组各维度的坐标网格 axis_grids = np.indices(original.shape) # 对坐标网格执行和原数组完全一致的切片操作 sliced_axis_grids = [grid[:, ::3] for grid in axis_grids] # 构造和切片数组同形状的索引数组,存储每个位置对应的原数组多维元组索引 indices_of_slice = np.empty(sliced_array.shape, dtype=object) for pos in np.ndindex(sliced_array.shape): indices_of_slice[pos] = tuple(ag[pos] for ag in sliced_axis_grids) # 校验逻辑(遍历对象数组需加refs_ok标记) for val, idx in zip(np.nditer(sliced_array), np.nditer(indices_of_slice, flags=["refs_ok"])): assert val == original[idx], "Error. This implementation is not correct. "
结果验证
打印两个数组的输出完全符合预期:
>>> print(sliced_array) [[5 3] [8 6]] >>> print(indices_of_slice) [[(0, 0) (0, 3)] [(1, 0) (1, 3)]]
实现原理
np.indices(original.shape)会生成和原数组维度一致的坐标网格,返回结果的第k个元素对应原数组第k个轴的所有位置坐标,形状和原数组完全相同- 对坐标网格应用和原数组完全相同的切片规则,即可得到切片覆盖区域内,每个元素在原数组各轴上的坐标值
- 遍历切片数组的所有位置,将对应位置的各轴坐标组合为元组存入对象数组,即可得到能和切片数组同步遍历的索引数组
- 该方法对任意维度的NumPy数组、任意合法切片规则都通用,只需要保证对坐标网格应用的切片逻辑和原数组完全一致即可
如果不需要适配
nditer逐个遍历的写法,不需要循环构造对象数组,直接执行indices_of_slice = np.stack(sliced_axis_grids, axis=-1)即可得到整数类型的索引数组,形状为(sliced_array.shape[0], sliced_array.shape[1], 2),通过indices_of_slice[i,j]即可拿到(i,j)位置对应的原数组坐标,性能更高。
内容的提问来源于stack exchange,提问作者Steinn Hauser Magnússon
相关产品推荐
相关产品推荐

