如何用线段树实现数组区间内大于P的最近值查询及更新操作
嘿,这个问题用线段树完全可以解决,而且能做到更新和查询都是O(log²N)的时间复杂度,咱们一步一步来拆解怎么做。
核心思路
要快速找到区间内大于P且与P差值最小的元素,本质上是找区间内大于P的最小元素。线段树的每个节点可以存储对应区间内的元素集合,并且保持有序——这样我们就能用二分查找快速定位到符合条件的元素。
线段树节点设计
每个线段树节点对应数组的一个子区间,节点内部存储该区间所有元素的有序集合:
- 如果是小规模数据,用普通的有序列表就行;
- 如果是大规模数据(比如N≥1e5),建议用支持O(logk)插入/删除/二分的有序集合(比如Python的
SortedList、C++的std::set),这样能保证操作效率。
构建线段树
- 叶子节点:对应数组的单个元素,有序集合里就只包含这一个元素。
- 非叶子节点:把左右子节点的有序集合合并成一个新的有序集合(因为两个子集合本身就是有序的,用双指针合并的时间是O(k),k是两个子集合的总长度)。
整个构建过程的时间复杂度是O(N logN)——每个元素会出现在logN个节点中,每层节点的合并总长度是N,总共logN层。
处理更新操作(类型1 U P)
当需要把数组第U个元素更新为P时:
- 找到对应位置的叶子节点,删除旧值,插入新值。
- 向上回溯到根节点,依次更新所有包含该位置的父节点:每个父节点都要删除旧值,插入新值,维持有序性。
如果用SortedList这类结构,每个节点的更新操作是O(logk)(k是节点内元素个数),从叶子到根有logN层,所以单次更新的时间复杂度是O(log²N)。
处理查询操作(类型2 L R P)
要找区间[L,R]内大于P的最小元素,我们可以把查询区间拆分成线段树中若干个不重叠的节点区间,然后逐个处理:
- 对每个拆分出的节点,用二分查找找到第一个大于P的元素(比如用
bisect_right)。 - 收集所有符合条件的候选元素,取其中最小的那个就是答案;如果没有候选元素,返回-1。
每个节点的二分查找是O(logk),拆分的节点数是O(logN),所以单次查询的时间复杂度也是O(log²N)。
代码实现(Python示例)
这里用sortedcontainers库的SortedList来保证高效操作,如果没有这个库,可以用bisect模块维护普通有序列表(适合小规模数据):
from sortedcontainers import SortedList class SegmentTree: def __init__(self, data): self.n = len(data) self.size = 1 # 找到大于等于n的最小2的幂,作为线段树的叶子节点数 while self.size < self.n: self.size <<= 1 # 初始化线段树,每个节点是一个SortedList self.tree = [SortedList() for _ in range(2 * self.size)] # 填充叶子节点 for i in range(self.n): self.tree[self.size + i].add(data[i]) # 构建上层节点 for i in range(self.size - 1, 0, -1): self.tree[i].update(self.tree[2 * i]) self.tree[i].update(self.tree[2 * i + 1]) def update_val(self, pos, new_val): # pos是题目中的1-based索引,转换成线段树的叶子节点索引 leaf_idx = self.size + pos - 1 # 获取旧值 old_val = self.tree[leaf_idx][0] # 更新叶子节点 self.tree[leaf_idx].discard(old_val) self.tree[leaf_idx].add(new_val) # 向上更新父节点 leaf_idx >>= 1 while leaf_idx >= 1: self.tree[leaf_idx].discard(old_val) self.tree[leaf_idx].add(new_val) leaf_idx >>= 1 def query_min_gt(self, l, r, p): # l和r是题目中的1-based区间 res = None # 转换成线段树的节点索引(1-based区间转内部索引) l_idx = self.size + l - 1 r_idx = self.size + r - 1 while l_idx <= r_idx: # 处理左节点(如果是右子节点) if l_idx % 2 == 1: lst = self.tree[l_idx] # 找到第一个大于p的元素的位置 idx = lst.bisect_right(p) if idx < len(lst): candidate = lst[idx] if res is None or candidate < res: res = candidate l_idx += 1 # 处理右节点(如果是左子节点) if r_idx % 2 == 0: lst = self.tree[r_idx] idx = lst.bisect_right(p) if idx < len(lst): candidate = lst[idx] if res is None or candidate < res: res = candidate r_idx -= 1 # 向上移动 l_idx >>= 1 r_idx >>= 1 return res if res is not None else -1 # 测试示例输入 if __name__ == "__main__": n = int(input()) data = list(map(int, input().split())) st = SegmentTree(data) q = int(input()) for _ in range(q): parts = input().split() if parts[0] == '1': u = int(parts[1]) p = int(parts[2]) st.update_val(u, p) else: l = int(parts[1]) r = int(parts[2]) p = int(parts[3]) print(st.query_min_gt(l, r, p))
复杂度分析
- 构建:O(N logN),每个元素在logN个节点中出现,每层合并总长度为N。
- 更新:O(log²N),每层节点的插入/删除是O(logk),共logN层。
- 查询:O(log²N),每个拆分节点的二分查找是O(logk),共logN个节点。
内容的提问来源于stack exchange,提问作者Rohit
相关产品推荐
相关产品推荐

