如何以Pythonic方式实现NumPy多维数组的元组多索引?
Numpy多维数组多索引的Pythonic实现
问题场景
我有一个多维numpy数组,已知其前N维和后M维的形状,示例如下:
>>> n = (3,4,5) >>> m = (6,) >>> a = np.ones(n + m) >>> a.shape (3, 4, 5, 6)
使用元组作为索引可以快速定位前N维的单个位置,返回后M维的结果:
>>> i = (1,1,2) >>> a[i].shape (6,)
但用列表作为索引无法达到相同效果:
>>> i = [1,1,2] >>> a[i].shape (3, 4, 5, 6)
当需要同时使用多个这样的元组索引时(比如提取多个前N维位置对应的后M维数据),常规写法无法得到预期结果:
>>> i = (1,1,2) >>> j = (2,2,2) # 期望得到形状为(2,6)的结果 >>> a[[i, j]].shape (2, 3, 4, 5, 6) # 实际结果不符合预期 >>> a[(i, j)].shape (3, 5, 6) # 同样不符合预期
需求适用于任意数量的索引,比如同时处理i、j、k等多个索引元组。
解决方案
利用numpy的花式索引特性,将多个索引元组转换为按维度分组的索引结构即可实现需求:
实现代码
import numpy as np n = (3,4,5) m = (6,) a = np.ones(n + m) # 定义多个索引元组 i = (1,1,2) j = (2,2,2) k = (0,0,0) # 将索引元组转置,得到按维度分组的索引数组 indices = tuple(np.array([i, j, k]).T) # 提取目标数据 result = a[indices] print(result.shape) # 输出: (3, 6) # 赋值操作同样适用 a[indices] = np.zeros((3,6))
原理说明
numpy的花式索引规则中,当索引是一个包含多个数组的元组时,每个数组对应原数组的一个维度。我们将多个索引元组堆叠成数组后转置,相当于把每个维度对应的所有索引值单独提取出来,形成(维度1索引数组, 维度2索引数组, ..., 维度N索引数组)的结构,numpy会自动广播这些索引,取出每个维度对应位置的元素组合,最终得到形状为(索引数量, M)的结果,完全符合需求。
这种方式无需循环,直接利用numpy原生机制,是更Pythonic的实现方式。
内容的提问来源于stack exchange,提问作者Simon
相关产品推荐
相关产品推荐

