线段树(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]'的线段树,其中的叶子节点表示的是原数组的元素,而非叶子节点表示的是原数组的部分区间。

segment tree

下面是线段树的一些基本操作和代码实现。

线段树的构建

线段树的构建可以使用递归的方式进行,每次递归处理区间'[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是当前节点的编号,lr是当前节点表示的区间的左右端点。

线段树的查询

线段树的查询可以使用递归的方式进行,每次递归处理区间'[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是当前节点的编号,lr是当前节点表示的区间的左右端点,qlqr是目标区间的左右端点。

线段树的修改

线段树的修改也可以使用递归的方式进行,每次递归处理区间'[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是当前节点的编号,lr是当前节点表示的区间的左右端点,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 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录