如何对NumPy数组按行索引,每行选取指定索引元素?
更快的NumPy行索引选取方法
当然有更高效的实现方式!你现在用的列表推导式本质是在Python层面逐行循环,当数组规模很大时,这种方式会因为Python循环的开销变得很慢——而NumPy的核心优势就是向量化运算,能把循环放到底层C代码中执行,速度会提升几个数量级。
下面给你两种最常用的高效实现方案:
方案1:利用高级索引(最直观推荐)
NumPy支持同时用行索引数组和列索引数组进行批量取值。我们只需要构造一个和ixs形状匹配的行索引数组,然后直接索引即可:
import numpy as np # 示例数据 X = np.array([[0,1,2,3], [4,5,6,7], [8,9,10,11]]) ixs = np.array([[1,3], [0,1], [2,0]]) # 构造行索引:每行对应X的行号,形状变为(n_rows, 1)以匹配ixs的列数 row_indices = np.arange(X.shape[0])[:, np.newaxis] # 批量取值 output = X[row_indices, ixs] print(output) # 输出: # [[ 1 3] # [ 4 5] # [10 8]]
这个方法完全利用了NumPy的向量化机制,没有任何Python循环,处理百万级别的数组时速度会比原方法快几十甚至上百倍。
方案2:使用np.take(扁平化索引方式)
另一种思路是将二维索引转换成扁平化的全局索引,然后用np.take批量取值:
# 计算每个元素的扁平化索引:行号 * 列数 + 列索引 flat_indices = ixs + np.arange(X.shape[0])[:, np.newaxis] * X.shape[1] # 批量取值并重塑形状 output = X.take(flat_indices).reshape(-1, 2) print(output) # 输出和方案1一致
这个方法的效率和方案1基本相当,只是实现思路不同,你可以根据自己的习惯选择。
性能对比示例
我们用百万行的数组来测试三种方法的速度:
import timeit # 生成大规模测试数据 X = np.random.rand(1_000_000, 10) ixs = np.random.randint(0, 10, (1_000_000, 2)) # 原方法 def original(): return np.array([row[ix] for row, ix in zip(X, ixs)]) # 方案1 def vectorized_1(): row_ix = np.arange(X.shape[0])[:, None] return X[row_ix, ixs] # 方案2 def vectorized_2(): flat_idx = ixs + np.arange(X.shape[0])[:, None] * X.shape[1] return X.take(flat_idx).reshape(-1, 2) # 测试执行时间(各运行1次) print(f"原方法耗时: {timeit.timeit(original, number=1):.2f} 秒") print(f"方案1耗时: {timeit.timeit(vectorized_1, number=1):.4f} 秒") print(f"方案2耗时: {timeit.timeit(vectorized_2, number=1):.4f} 秒")
运行结果大概是这样(取决于机器性能):
原方法耗时: 2.15 秒
方案1耗时: 0.0312 秒
方案2耗时: 0.0289 秒
可以看到向量化方法的速度提升非常明显!
内容的提问来源于stack exchange,提问作者mxbi
相关产品推荐
相关产品推荐

