使用Numpy函数按索引向量从多维数组提取值的优雅实现方案问询
优雅提取目标数组的Numpy方法
嘿,这个问题问得很到位!你当前的写法虽然能运行,但其实存在一个小问题:[True] * len(vec)会选中第三维度的所有元素,得到的数组形状应该是(4,5,4),而不是你想要的(4,5)。不过没关系,我们可以用Numpy的高级索引特性,写出更简洁且符合需求的代码~
核心思路
你的目标本质是:对每个索引位置i(从0到3),提取mat[vec[i], :, i](即第一维度取vec[i]、第二维度全选、第三维度取第i个元素),然后将这些长度为5的一维数组堆叠成(4,5)的二维数组。
优雅解决方案
直接用np.arange生成第三维度的索引序列,和vec配合进行高级索引:
import numpy as np # 示例数据 mat = np.random.rand(3, 5, 4) # shape (3,5,4) vec = np.random.randint(0, 3, size=4) # shape (4,),元素0-2 # 优雅提取目标数组 result = mat[vec, :, np.arange(vec.size)]
这段代码的效果完全符合你的需求:vec是第一维度的索引数组(shape (4,)),np.arange(vec.size)是第三维度的索引数组(shape (4,)),Numpy会自动将这两个数组广播匹配,对每个i提取mat[vec[i], :, i],最终拼接成(4,5)的数组。
为什么这更优雅?
- 无需手动创建布尔列表,完全利用Numpy原生的索引机制,代码更简洁易读
- 明确表达了“第三维度取对应位置元素”的意图,逻辑更清晰
- 性能上和手动创建布尔列表相当,但避免了额外的列表创建开销
另一种等价写法(更简洁)
如果觉得np.arange(vec.size)有点长,也可以用np.arange(len(vec)),甚至直接用range(len(vec))(Numpy会自动将其转换为数组):
result = mat[vec, :, range(len(vec))]
内容的提问来源于stack exchange,提问作者Solaris
相关产品推荐
相关产品推荐

