Python中如何高效将多区间判断if语句转换为数组索引实现
高效实现升序不重叠区间的索引匹配
背景:布尔索引的简化写法
如下是一段基础的奇偶判断代码:
def IsEven(n): if n%2==0: return "Is even" else: return "Is odd"
可以利用Python里布尔值可作为整数索引的特性,直接简化为通过数组索引取值的形式:
def isEven(n): return ["Is even","Is odd"][n%2==0]
多区间判断的改造需求
我现在有一段多区间判断的代码:
def intervalsToOutput(n): intervals=[(x1,x2), (x3,x4), (x5,x6), ... (xn,xn+1)] if x1<=n<=x2: return "in first interval" elif x3<=n<=x4: return "in second interval" elif x5<=n<=x6: return "in third interval" ... elif xn<=n<=xn+1: return "in last interval"
已知所有区间$(x_i,x_{i+1})$互不重叠,且按区间起点$x_i$升序排列,能不能把它高效替换成「匹配n所属区间的索引→直接读取答案数组对应值」的形式,目标实现结构如下:
def intervalsToOutput(n): intervals=[(x1,x2), (x3,x4), (x5,x6), ... (xn,xn+1)] answer=["in first interval","in second interval","in third interval",...,"in last interval"] return answer[index of (n in Interval)]
现有bisect方案的问题
我目前以运行速度为优先目标写出的最优方案,用了标准库的bisect模块做二分查找:
def intervalsToOutput(n): intervals=[(x1,x2), (x3,x4), (x5,x6), ... (xn,xn+1)] answer=["in first interval","in second interval","in third interval",...,"in last interval"] import bisect as bisect return answer[bisect.bisect_left(intervals, (n, )) - 1]# 零基列表所以索引减1
bisect_left本身能以$O(\log n)$的时间复杂度查找元素插入位置,逻辑是将元组(n,)和存储区间的元组做比较,但它的设计目标不是做区间匹配:当n落在两个相邻区间的间隙中(满足$x_{i+1}<n\leq x_j$)时,返回的索引是错误的,无法正确匹配所属区间,间隙位置示意如下:
... (x_i, x_{i+1}), (x_j, x_{j+1}) ... *[x_i______x_i+1]* 该区间后间隙位置的n会匹配失败 x_j *[________x_j+1]*
尝试过的intervaltree方案
补充说明:我也实现了用intervaltree库替代bisect的方案,但这个方案有额外开销:intervaltree的查询返回值是集合,结果还可能为空。我也接受基于这段代码优化的更快速、更简洁的解决方案,具体代码如下:
# 生成区间和对应答案 import random min=0 max=20 numIntervals=6 def Intervals(min,max,n): randInts = random.sample(range(min, max), n * 2) randInts.sort() return [(x1,x2) for x1,x2 in zip(randInts[::2], randInts[1::2])] intervals=Intervals(min,max,numIntervals) answer=[f"In {n}th interval" for n in range(len(intervals))] # 基于intervaltree的实现 import intervaltree tree = intervaltree.IntervalTree() [tree.addi(i[0],i[1],a) for i,a in zip(intervals,answer)] def intervalsToOutput(x): res=tree.at(x) if len(res)>0: return res.pop().data return "value not found" print(intervals) [print(f"{x} is ",intervalsToOutput(x)) for x in random.sample(range(min,max),numIntervals)]
内容的提问来源于stack exchange,提问作者Colim
相关产品推荐
相关产品推荐

