如何利用numpy数组arr1作为列索引快速提取arr2对应元素生成arr3?
高效从NumPy数组中按行索引提取元素的方法
你当前用列表推导式的方法虽然能得到正确结果,但在处理大数据量时会因为Python层面的循环产生不小的性能开销。NumPy提供了向量化的高级索引方案,能直接在底层完成操作,效率提升非常明显。
最优实现代码
import numpy as np # 假设arr1和arr2已定义 arr3 = arr2[np.arange(arr2.shape[0]), arr1]
原理说明
np.arange(arr2.shape[0])生成和arr2行数一致的行索引数组(比如arr2有N行,就生成[0,1,2,...,N-1])- 将行索引数组和arr1(列索引数组)组合,NumPy会自动按
(行索引[i], 列索引[i])的位置提取元素,直接生成目标数组arr3
性能对比示例
用百万级数据测试两种方法的耗时:
# 生成测试数据 arr1 = np.random.randint(0, 2, size=1_000_000) arr2 = np.random.rand(1_000_000, 2).astype(np.float32) # 列表推导式方法 %timeit np.array([arr2[i][arr1[i]] for i in range(len(arr2))]) # 输出示例:1.2 s ± 45.2 ms per loop (mean ± std. dev. of 7 runs, 1 loop each) # 向量化索引方法 %timeit arr2[np.arange(arr2.shape[0]), arr1] # 输出示例:2.1 ms ± 101 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
可以看到,向量化方法的速度是列表推导式的几百倍,数据量越大,优势越明显。
额外补充
如果你的arr1是从0开始的合法列索引,还可以用另一种写法(效果完全一致):
arr3 = arr2[np.arange(len(arr2)), arr1]
两种写法的核心都是利用NumPy的高级索引机制,避免Python循环的开销。
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

