如何用一维np.ndarray索引高效提取二维np.ndarray对应元素?
问题:高效提取NumPy数组指定索引元素
现有两个np.ndarray数组:二维数组x,一维数组y,其中y[i]为要从x[i]子数组中提取的元素索引。当前用Python列表推导式实现了需求,但不够简洁优雅,且对执行速度要求极高,生成二维掩码的方案不可行,求更优实现方式。
示例代码:
import numpy as np x = np.array([[.55, .45], [0.78, .22], [.85, .15]]) y = np.array([1,0,1]) preds = np.array([x[i, y[i]] for i in range(y.shape[0])]) print(preds) #[0.45, 0.78, 0.15] <- 0.45 == x[0][1], 0.78 == x[1][0], 0.15 == x[2][1]
最优解决方案
用NumPy的高级索引直接实现,这是完全向量化的操作,比列表推导式快得多,尤其在数组规模较大时优势明显。
核心思路是用np.arange(x.shape[0])生成行索引数组,和y的列索引数组配对,直接从x中提取对应位置的元素:
import numpy as np x = np.array([[.55, .45], [0.78, .22], [.85, .15]]) y = np.array([1,0,1]) # 向量化提取 preds = x[np.arange(x.shape[0]), y] print(preds) # 输出: [0.45 0.78 0.15]
方案优势
- 高效:避免Python级循环,所有操作在NumPy底层C代码执行,速度远超列表推导式。
- 简洁:一行代码完成需求,可读性更强。
- 内存友好:无需生成额外掩码数组,直接通过索引定位元素,内存占用更低。
内容的提问来源于stack exchange,提问作者Aerocobra
相关产品推荐
相关产品推荐

