如何通过另一numpy数组的索引高效索引numpy数组?
高效实现numpy数组的指定索引取值
原生numpy最优方案:高级索引直接完成
不需要扩展L的形状,也不用扁平化数组,利用numpy的广播式高级索引就能高效实现需求,完全避免中间数组的额外开销:
import numpy as np # 示例数据 a, b, c = 3, 4, 5 A = np.random.rand(a, b, c) L = np.random.randint(0, a, size=(a, b)) # 核心实现 B = A[L, np.arange(b)[np.newaxis, :], :]
原理说明
L是形状为(a,b)的数组,作为第一个维度的索引,numpy会自动将其与后续索引广播匹配np.arange(b)[np.newaxis, :]生成形状为(1,b)的数组,对应第二个维度的j索引,广播后变为(a,b),和L的形状完全对齐- 第三个维度用
:直接取所有元素,最终得到的B形状为(a,b,c),完全符合需求
使用einops的简洁实现
如果你习惯用einops,也可以借助其维度重排语法让逻辑更直观,本质还是基于numpy的索引机制:
from einops import rearrange # 调整维度顺序后索引,再还原 A_jik = rearrange(A, 'i j k -> j i k') B = rearrange(A_jik[np.arange(b), L.T, :], 'j i k -> i j k')
或者更紧凑的一行写法:
B = rearrange(A[L, np.arange(b)[None, :], :], 'i j k -> i j k')
对比扁平化方法的优势
你之前用的扁平化+重塑形方法虽然可行,但完全没必要——numpy原生的高级索引已经支持这种多维对齐的索引逻辑,不需要额外创建中间数组,内存占用和运行效率都更优,写法也更简洁直接。
内容的提问来源于stack exchange,提问作者Danny Duberstein
相关产品推荐
相关产品推荐

