线段树:区间查询利器,代码讲解与实现
线段树(Segment Tree)是一种用来解决区间查询问题(如区间最小值、区间和等)的数据结构。它可以支持区间修改和单点查询等操作,时间复杂度为$O(log\ n)$。
线段树是一棵二叉树,每个节点表示区间'[l,r]',其中l和r是区间的左右端点。每个节点有两个子节点,分别表示区间'[l,mid]'和'[mid+1,r]',其中$mid=\left\lfloor\frac{l+r}{2}\right\rfloor$。线段树的叶子节点表示区间的单个元素,而非叶子节点表示的区间是其子节点表示的区间的并集。例如,下图展示了一个表示区间'[1,8]'的线段树,其中的叶子节点表示的是原数组的元素,而非叶子节点表示的是原数组的部分区间。
下面是线段树的一些基本操作和代码实现。
线段树的构建
线段树的构建可以使用递归的方式进行,每次递归处理区间'[l,r]'的左右两个子区间,直到区间被划分成单个元素。
def build_tree(a, tree, node, l, r):
if l == r:
tree[node] = a[l]
else:
mid = (l + r) // 2
build_tree(a, tree, node * 2, l, mid)
build_tree(a, tree, node * 2 + 1, mid + 1, r)
tree[node] = tree[node * 2] + tree[node * 2 + 1]
其中,a是原数组,tree是线段树的数组,node是当前节点的编号,l和r是当前节点表示的区间的左右端点。
线段树的查询
线段树的查询可以使用递归的方式进行,每次递归处理区间'[l,r]'的左右两个子区间,直到目标区间被完全覆盖或不被覆盖。
def query_tree(tree, node, l, r, ql, qr):
if qr < l or ql > r: # 不相交
return 0
elif ql <= l and qr >= r: # 完全包含
return tree[node]
else: # 相交
mid = (l + r) // 2
left_sum = query_tree(tree, node * 2, l, mid, ql, qr)
right_sum = query_tree(tree, node * 2 + 1, mid + 1, r, ql, qr)
return left_sum + right_sum
其中,tree是线段树的数组,node是当前节点的编号,l和r是当前节点表示的区间的左右端点,ql和qr是目标区间的左右端点。
线段树的修改
线段树的修改也可以使用递归的方式进行,每次递归处理区间'[l,r]'的左右两个子区间,直到目标区间被完全覆盖或不被覆盖。在修改时,需要注意修改的值是原来的值加上修改的值,而非直接覆盖原有值。
def update_tree(tree, node, l, r, idx, val):
if l == r:
tree[node] += val
else:
mid = (l + r) // 2
if idx <= mid:
update_tree(tree, node * 2, l, mid, idx, val)
else:
update_tree(tree, node * 2 + 1, mid + 1, r, idx, val)
tree[node] = tree[node * 2] + tree[node * 2 + 1]
其中,tree是线段树的数组,node是当前节点的编号,l和r是当前节点表示的区间的左右端点,idx是要修改的元素在原数组中的下标,val是要加上的值。
完整代码
下面是一个实现线段树的完整代码,包括线段树的构建、查询和修改操作。
def build_tree(a, tree, node, l, r):
if l == r:
tree[node] = a[l]
else:
mid = (l + r) // 2
build_tree(a, tree, node * 2, l, mid)
build_tree(a, tree, node * 2 + 1, mid + 1, r)
tree[node] = tree[node * 2] + tree[node * 2 + 1]
def query_tree(tree, node, l, r, ql, qr):
if qr < l or ql > r: # 不相交
return 0
elif ql <= l and qr >= r: # 完全包含
return tree[node]
else: # 相交
mid = (l + r) // 2
left_sum = query_tree(tree, node * 2, l, mid, ql, qr)
right_sum = query_tree(tree, node * 2 + 1, mid + 1, r, ql, qr)
return left_sum + right_sum
def update_tree(tree, node, l, r, idx, val):
if l == r:
tree[node] += val
else:
mid = (l + r) // 2
if idx <= mid:
update_tree(tree, node * 2, l, mid, idx, val)
else:
update_tree(tree, node * 2 + 1, mid + 1, r, idx, val)
tree[node] = tree[node * 2] + tree[node * 2 + 1]
if __name__ == '__main__':
a = [1, 3, 5, 7, 9, 11]
n = len(a)
tree = [0] * (4 * n)
build_tree(a, tree, 1, 0, n - 1)
print(query_tree(tree, 1, 0, n - 1, 0, 2)) # 9
update_tree(tree, 1, 0, n - 1, 1, 2)
print(query_tree(tree, 1, 0, n - 1, 0, 2)) # 11
原文地址: https://www.cveoy.top/t/topic/nGDJ 著作权归作者所有。请勿转载和采集!