如何高效对整数数组执行多组区间比较统计?
问题描述
我有一个整数数组:
import numpy as np mypos = np.array([10, 20, 30, 40, 50])
以及一组由起始和结束位置组成的区间列表:
mydelims = [[5, 12], [15,31], [12,16], [22,69]]
我想要遍历mydelims,统计每个区间内包含的数组元素数量。一开始写了这段代码:
for mypair in mydelims: print(sum(mypos>mypair[0] & mypos<mypair[1]))
但Python无法正常运行。已知mypos>42会返回布尔数组,想知道对整数数组执行多组区间比较的最高效方法是什么?
错误原因与修正
你写的代码报错核心是运算符优先级问题:&(按位与)的优先级比比较运算符>、<更高,导致代码实际执行的是mypair[0] & mypos,完全偏离了原本的区间判断逻辑。
只需要给每个比较表达式加上括号,修正后的代码就能正常运行:
for mypair in mydelims: print(sum((mypos > mypair[0]) & (mypos < mypair[1])))
执行后会依次输出每个区间内的元素数量:1、2、0、3。
高效实现方法
如果处理的数组或区间数量较大,Python层面的循环效率较低,推荐用NumPy的矢量化操作来提升性能:
方法一:利用np.searchsorted(最优大数量场景)
searchsorted基于二分查找实现,能快速找到边界在有序数组中的插入位置,通过左右位置的差值得到区间内元素数量,时间复杂度为O(n log m)(n为区间数,m为数组长度),适合大数据量场景:
import numpy as np mypos = np.array([10, 20, 30, 40, 50]) mydelims = np.array([[5, 12], [15,31], [12,16], [22,69]]) # 若mypos未排序,先执行mypos.sort() left_idx = np.searchsorted(mypos, mydelims[:, 0], side='right') right_idx = np.searchsorted(mypos, mydelims[:, 1], side='left') counts = right_idx - left_idx print(counts) # 输出: [1 2 0 3]
side='right'表示左边界的插入位置取右侧,确保只统计大于左边界的元素side='left'表示右边界的插入位置取左侧,确保只统计小于右边界的元素
方法二:矢量化广播比较
利用NumPy的广播机制,将区间数组和元素数组扩展为二维布尔矩阵,再按行求和得到每个区间的元素数量,适合中等数据量场景:
import numpy as np mypos = np.array([10, 20, 30, 40, 50]) mydelims = np.array([[5, 12], [15,31], [12,16], [22,69]]) # 通过newaxis扩展维度,实现广播比较,得到(4,5)的布尔矩阵 mask = (mypos > mydelims[:, 0, np.newaxis]) & (mypos < mydelims[:, 1, np.newaxis]) counts = mask.sum(axis=1) print(counts) # 输出: [1 2 0 3]
这种方法把循环逻辑放到NumPy底层的C实现中执行,效率远高于Python循环。
内容的提问来源于stack exchange,提问作者Whitehot
相关产品推荐
相关产品推荐

