如何高效统计多个不等长有序numpy数组在指定区间的元素数?
Question
我现在有一组不等长的一维有序numpy数组(比如示例里的M0、M1、M2),需要统计每个数组在由另一有序数组zbin相邻元素定义的数值区间内的元素数量。
当前我已经实现了一个方法,但因为实际任务里zbin的长度、数组的规模和数量都远大于示例,而且所有数组都是有序的,所以想问问:当前的search函数和实现方案是不是最优最快的?有没有更高效的实现方式?
示例代码与输出
""" Function to do search query """ def search(numrange, lst): arr = np.zeros(len(lst)) for i in range(len(lst)): probe = lst[i] count = 0 for j in range(len(probe)): if (probe[j]>numrange[1]): break if (probe[j]>=numrange[0]) and (probe[j]<=numrange[1]): count = count + 1 arr[i] = count return arr """ Some example of sorted one-dimensional arrays of unequal lengths """ M0 = np.array([5.1, 5.4, 6.4, 6.8, 7.9]) M1 = np.array([5.2, 5.7, 8.8, 8.9, 9.1, 9.2]) M2 = np.array([6.1, 6.2, 6.5, 7.2]) """ Implementation and output """ lst = [M0, M1, M2] zbin = np.array([5.0, 5.5, 6.0, 6.5]) zarr = np.zeros( (len(zbin)-1, len(lst)) ) for i in range(len(zbin)-1): numrange = [zbin[i], zbin[i+1]] zarr[i,:] = search(numrange, lst) print(zarr)
输出:
[[ 2. 1. 0.] [ 0. 1. 0.] [ 1. 0. 3.]]
输出的zarr每行对应zbin的一个区间(比如第一行对应[5.0,5.5]),每列对应一个数组的统计数。比如[5.0,5.5]区间里,M0有2个元素,M1有1个,M2有0个。
Answer
你的当前实现虽然聪明地利用了数组有序的特性(遇到大于区间上限就break),但双重循环的结构在处理大数据量时肯定会拖慢速度——尤其是当数组数量多、每个数组元素量大的时候,Python层面的循环天生就是性能瓶颈。
既然所有数组都是有序的,那必须用上numpy的np.searchsorted啊!这个函数专门针对有序数组做快速二分查找,时间复杂度是O(log n)每次查询,比你手动遍历的O(n)快得多,而且底层是C实现的,效率拉满。
优化后的实现方案
核心逻辑很简单:对每个有序数组,用searchsorted找到zbin里每个分界点在数组中的插入位置,然后相邻位置的差值就是对应区间内的元素数量——完全不需要遍历每个元素!
直接看代码:
import numpy as np # 示例数据 M0 = np.array([5.1, 5.4, 6.4, 6.8, 7.9]) M1 = np.array([5.2, 5.7, 8.8, 8.9, 9.1, 9.2]) M2 = np.array([6.1, 6.2, 6.5, 7.2]) lst = [M0, M1, M2] zbin = np.array([5.0, 5.5, 6.0, 6.5]) # 高效统计函数 def count_in_bins(sorted_arrays, bins): # 初始化结果数组:行数=区间数,列数=数组数 result = np.zeros((len(bins)-1, len(sorted_arrays)), dtype=int) for idx, arr in enumerate(sorted_arrays): # 找到每个bin分界点在arr中的插入位置(side='right'确保包含等于上限的元素) positions = np.searchsorted(arr, bins, side='right') # 相邻位置相减,得到每个区间的元素数量 result[:, idx] = positions[1:] - positions[:-1] return result # 执行并输出 zarr = count_in_bins(lst, zbin) print(zarr)
输出和你的原代码完全一致:
[[2 1 0] [0 1 0] [1 0 3]]
为什么这个方案更高效?
- 二分查找替代手动遍历:
np.searchsorted的C实现比Python循环快几个数量级,元素越多,差距越明显。 - 一次性计算所有区间:对每个数组只需要一次查找,就能得到所有区间的统计数,不用像原代码那样对每个区间单独处理。
- 更少的Python层循环:只需要遍历每个数组一次,核心计算都在numpy的底层完成,避免了Python循环的额外开销。
如果你的数组数量特别多,还可以用列表推导来简化代码,但上面的写法已经足够高效了——毕竟核心逻辑都是向量化的。
内容的提问来源于stack exchange,提问作者Commoner
相关产品推荐
相关产品推荐

