Python中如何用低维索引实现numpy多维数组的向量化复杂索引
报错原因
你直接使用x[indLargest2ndAxis, 1]时,numpy会把indLargest2ndAxis作为数组第0轴的索引,而你得到的索引是对应第2轴的取值(范围0-9),数组第0轴长度仅为5,自然会出现索引越界的错误。
解决方案
方案1:使用np.take_along_axis(推荐)
这个API是numpy专门为沿指定轴、按对应索引提取数值设计的,代码简洁易读,输出自动匹配你需要的形状:
# 先提取最后一个轴第1位的所有数值,得到形状为(5,10,10)的数组 target_arr = x[..., 1] # 沿第2轴按indLargest2ndAxis的索引提取数值,输出形状为(5,10,6) res = np.take_along_axis(target_arr, indLargest2ndAxis, axis=2)
方案2:手动构造广播索引
如果你需要更灵活的自定义索引逻辑,可以手动构造前两个轴的广播索引,配合你得到的第2轴索引使用:
# 构造第0、1轴的基础索引,形状分别为(5,1)、(1,10),可自动广播匹配到(5,10,6) i, j = np.ogrid[:x.shape[0], :x.shape[1]] res = x[i, j, indLargest2ndAxis, 1]
结果验证
两种方案的输出都和你给出的循环实现完全一致,可通过如下代码验证:
loop_res = np.array([ [ [ x[i, j, k, 1] for k in indLargest2ndAxis[i,j]] for j in range(10) ] for i in range(5) ]) print(np.allclose(res, loop_res)) # 输出为True
内容的提问来源于stack exchange,提问作者Puco4
相关产品推荐
相关产品推荐

