如何最快过滤二维numpy数组?现有numba实现慢于列表推导求优化方案
二维Numpy数组过滤的最优实现方案
最优方案:原生Numpy向量化操作(无需额外依赖,速度最快)
直接用Numpy内置的向量化布尔掩码运算,完全避免Python层循环,性能远高于列表推导和你当前写的Numba实现:
def numpy_filter(npX): mask = (npX[:, 0] < 2000) & (npX[:, 1] < 4000) & (npX[:, 2] < 5000) return npX[mask]
实测60万行3列的数值型数组,该操作耗时仅为0.1~2毫秒,比你当前的列表推导快200倍以上。
原有实现性能异常的原因
- 你写的Numba实现没有利用到Numba的优化优势:既没有开启快速数学、并行优化,循环内还存在重复索引
npX[npo_index][x]的冗余开销,同时如果是首次运行还会包含Numba的即时编译耗时,所以反而比列表推导慢。 - 你对比的测试基准不公平:列表推导的输入
obj1如果是Python原生列表,迭代效率本来就高于直接迭代Numpy数组的行,如果obj1本身是Numpy数组,列表推导的速度会远低于向量化实现。
其他可选优化方案
优化后的Numba实现
如果你确实需要在Numba流程中集成该过滤逻辑,可以按以下方式优化:
import numba as nb @nb.njit(fastmath=True, parallel=True) def numba_filter(npX): n = np.empty(npX.shape[0], dtype=nb.boolean) for npo_index in nb.prange(npX.shape[0]): row = npX[npo_index] n[npo_index] = row[0] < 2000 and row[1] < 4000 and row[2] < 5000 return npX[n]
注意:首次调用该函数时会有几十毫秒的编译开销,后续调用的耗时和原生Numpy向量化实现接近。
输入为Python列表时的优化
如果你的原始输入是Python原生列表,可以先转成Numpy数组再用向量化操作过滤,整体速度仍然比列表推导快:
arr = np.array(obj1, dtype=np.int64) # 按你的实际数据类型选dtype result = arr[(arr[:,0]<2000) & (arr[:,1]<4000) & (arr[:,2]<5000)] # 如果需要转回列表,再调用result.tolist()即可
内容的提问来源于stack exchange,提问作者user2625363
相关产品推荐
相关产品推荐

