CC++ & Algorithm

线段树的区间查询

极难2
语言版本:通用
概述:学习如何在线段树中快速查询任意区间内的统计信息,利用递归覆盖和剪枝实现O(log n)复杂度。

什么是线段树的区间查询?

想象你有一本全班同学的考试成绩册,你想快速知道第2号到第6号同学的总分。如果你一个一个加,要加5个数,但如果有1000个同学,你就得加很多次。线段树就像提前把成绩分小组统计好,比如1-4号一组、5-8号一组,每组存好了总分。查询时,你只需要看哪些小组完全包含在你的查询范围里,把它们的总分加起来,而不用重新计算每个同学。这种快速“区间查询”就是线段树最拿手的本领。

区间查询:给定一个区间 [L, R](通常从1开始编号),求原数组中这个区间内所有数的和(或其他统计量,比如最大值、最小值、最大值出现次数等)。线段树能在 O(log n) 时间内完成查询,比直接遍历 O(n) 快很多。


一、核心思想:分治 + 剪枝

线段树已经将数组分成一层层的小区间(节点)。查询时,我们从根节点(代表整个数组)出发,递归往下走,根据当前节点区间与查询区间 [L, R] 的关系,决定下一步:

  1. 完全覆盖:如果当前节点区间完全落在 [L, R] 内,那么直接返回这个节点存好的值(比如该区间的和、最大值等)。这是最理想的情况,不用继续往下走。
  2. 完全不相交:如果当前节点区间和 [L, R] 没有重叠,返回一个对结果没有影响的“空值”(求和时返回0,求最大值时返回一个很小的数,比如 -∞)。
  3. 部分重叠:否则,当前区间有一部分在 [L, R] 内,一部分在外面。这时需要继续递归到左右孩子节点,然后合并左右孩子返回的结果。

这个过程像“查地图”:要找某个区域,先看整个地图是否全在范围内?不是,就切一半;再分别看两半……直到找到完全在范围内的子区域。


二、生活中的例子:班级成绩分组统计

沿用之前的班级成绩分组管理。假设全班8个同学,学号1~8,成绩分别为:1, 3, 5, 7, 9, 11, 13, 15。我们已经构建好了线段树(每个节点存对应区间的和)。

现在想知道学号2到6号同学的总分。查询过程如下(对照后面代码理解):

  • 从根节点 [1,8] 开始:它太大,不完全包含 [2,6] → 部分重叠,往左右走。
  • 左孩子 [1,4]:部分重叠(因为 [2,6] 只包含 [2,4] 部分)→ 再往左右。
    • 左孙 [1,2]:部分重叠(只包含2号)→ 往下。
      • [1,1]:与 [2,6] 无交集 → 返回0。
      • [2,2]:完全包含 → 返回成绩3。
    • 右孙 [3,4]:完全包含(因为 [3,4] 完全在 [2,6] 内)→ 返回和 5+7=12。
  • 右孩子 [5,8]:部分重叠(只包含 [5,6] 部分)→ 往下。
    • 左孙 [5,6]:完全包含 → 返回和 9+11=20。
    • 右孙 [7,8]:无交集 → 返回0。

最终结果 = 3 + 12 + 20 = 35。验证:直接加 3+5+7+9+11 = 35,正确。

整个过程只访问了少数几个节点(红色框住的节点),不需要遍历所有8个元素。这就是线段树的高效之处。


三、查询函数详解(通用伪代码)

// 查询区间[L, R] 的统计值
int query(int p, int l, int r, int L, int R) {
    // 情况1:完全覆盖
    if (L <= l && r <= R) {
        return tree[p];
    }
    // 情况2:无交集
    if (r < L || l > R) {
        return 0;          // 求和返回0;求最大值返回-inf
    }
    // 情况3:部分重叠,继续递归
    int mid = (l + r) / 2;
    int left_res = query(p * 2, l, mid, L, R);
    int right_res = query(p * 2 + 1, mid + 1, r, L, R);
    return left_res + right_res;  // 合并(此处为加法)
}

关键点

  • p 是节点在数组中的下标,lr 是该节点负责的区间。
  • 完全覆盖的判断条件 L <= l && r <= R 一定要写对,不能把 lL 搞混。
  • 无交集的判断 r < L || l > R 要包含等号情况(区间端点不相交)。
  • 递归时分别传 p*2(左孩子)和 p*2+1(右孩子)。

四、常见错误与避坑指南

  1. 忘记处理无交集情况:如果不判断,程序会继续递归,可能陷入死循环或访问到越界节点。一定要在递归前检查。
  2. 完全覆盖条件写反:比如写成 l <= L && R <= r,这样会导致只有查询区间完全被当前区间包含时才返回,而我们想要的是“当前区间完全被查询区间包含”。一定要记住:从查询角度看,条件应该是查询区间包住了当前区间
  3. 递归传参错误:比如在递归左孩子时不小心传成了 (p*2, l, r, L, R),没有更新区间范围。必须传递正确的 l, midmid+1, r
  4. 数组越界:线段树大小通常是 4 * n,但如果你查询的 L,R 超出数组范围(比如 L=0),需要提前处理或保证输入合法。
  5. 合并操作与构建不一致:如果构建时用的是求和,查询时也必须求和;如果构建时存的是最大值,查询也要返回最大值合并(max(left, right))。不能混用。

五、完整可运行代码(C++ 版,带详细注释)

#include <iostream>
using namespace std;

const int MAXN = 10000;        // 最大数据规模
int a[MAXN];                   // 原始数组,下标从1开始
int tree[4 * MAXN];            // 线段树数组,大小设为4倍

// 构建线段树
void build(int p, int l, int r) {
    if (l == r) {               // 叶子节点,直接赋值
        tree[p] = a[l];
        return;
    }
    int mid = (l + r) / 2;      // 取中间点
    build(p * 2, l, mid);       // 递归构建左子树
    build(p * 2 + 1, mid + 1, r); // 递归构建右子树
    tree[p] = tree[p * 2] + tree[p * 2 + 1]; // 合并:求和
}

// 区间查询:返回区间[L,R]的和
int query(int p, int l, int r, int L, int R) {
    // 情况1:当前区间完全被查询区间覆盖
    if (L <= l && r <= R) {
        return tree[p];
    }
    // 情况2:当前区间与查询区间无交集
    if (r < L || l > R) {
        return 0;    // 求和返回0;求最大值返回-1e9
    }
    // 情况3:部分重叠,递归查询左右子树
    int mid = (l + r) / 2;
    int left_res = query(p * 2, l, mid, L, R);          // 左子树结果
    int right_res = query(p * 2 + 1, mid + 1, r, L, R); // 右子树结果
    return left_res + right_res;   // 合并结果(此处为加法)
}

int main() {
    int n = 8;
    // 初始化数组:1,3,5,7,9,11,13,15
    for (int i = 1; i <= n; i++) {
        a[i] = 2 * i - 1;
    }
    build(1, 1, n);   // 从根节点1开始,区间[1,n]

    int L = 2, R = 6; // 查询学号2到6
    int ans = query(1, 1, n, L, R);
    cout << "Sum of interval [" << L << "," << R << "] = " << ans << endl;
    // 输出:Sum of interval [2,6] = 35
    return 0;
}

Python 版本(类封装,对外接口清晰)

class SegmentTree:
    def __init__(self, data):
        self.n = len(data)              # 数组长度
        self.a = [0] + data             # 1-based 数组,前面补0
        self.tree = [0] * (4 * (self.n + 1))  # 线段树数组

    def build(self, p, l, r):
        """构建线段树,p为节点编号,l,r为区间范围"""
        if l == r:
            self.tree[p] = self.a[l]
            return
        mid = (l + r) // 2
        self.build(p * 2, l, mid)
        self.build(p * 2 + 1, mid + 1, r)
        self.tree[p] = self.tree[p * 2] + self.tree[p * 2 + 1]

    def query(self, p, l, r, L, R):
        """内部递归查询,返回区间[L,R]的和"""
        if L <= l and r <= R:            # 完全覆盖
            return self.tree[p]
        if r < L or l > R:              # 无交集
            return 0
        mid = (l + r) // 2
        left = self.query(p * 2, l, mid, L, R)
        right = self.query(p * 2 + 1, mid + 1, r, L, R)
        return left + right

    def range_query(self, L, R):
        """外部接口:查询区间[L,R]的和(1-based)"""
        return self.query(1, 1, self.n, L, R)

# 测试
if __name__ == "__main__":
    arr = [1, 3, 5, 7, 9, 11, 13, 15]
    st = SegmentTree(arr)
    st.build(1, 1, st.n)
    print("Sum [2,6] =", st.range_query(2, 6))  # 输出:35

六、不止求和:其他统计量的查询方法

线段树的查询思想通用,只需修改合并操作:

统计量完全覆盖时返回无交集时返回合并操作
求和tree[p]0left + right
最大值tree[p]-INF (比如-1e9)max(left, right)
最小值tree[p]INF (比如1e9)min(left, right)
最大公约数tree[p]0gcd(left, right)
区间乘积tree[p]1left × right

例子:如果我们要查询 [L, R] 的最大值,只需在 query 中:

  • 无交集时返回一个非常小的数(如 -1e9),保证不影响最大值。
  • 合并时取 max(left_res, right_res)

注意:构建时也要对应修改 tree[p] 的合并方式(比如 tree[p] = max(tree[p*2], tree[p*2+1]))。


七、性能小贴士

  • 线段树查询的时间复杂度是 O(log n),因为每次递归深度不超过树高,树高约为 log₂ n。
  • 实际访问节点数大约是 4*log₂ n(因为部分重叠会触发两路递归),但远小于 n。
  • 对比直接遍历:n=10万时,直接遍历要10万步,线段树只需要约20步左右,快几千倍。

八、相关指引

学完区间查询后,下一步可以学习:

  • 单点更新:如何在线段树中修改一个元素的值,同时更新所有相关祖先节点的信息(也是 O(log n))。
  • 区间更新与懒标记:如果要一次性给整个区间加上同一个数,如何高效完成(引入延迟标记,避免递归到叶子)。
  • 线段树的其他应用:比如求区间最大子段和、区间极差等。

线段树是很多复杂数据结构的基础,掌握好区间查询,你就已经会了它最常用的操作。试着用上面的代码查询一下你自己班级某次考试的总分或最高分吧!

例题精讲

1单选题

线段树中,区间查询的时间复杂度是多少?

AO(1)
BO(log n)
CO(n)
DO(n log n)
2单选题

假设有一个长度为 8 的数组,构建了一棵线段树。若要查询区间 [3,6] 的和,以下哪项描述是正确的?

A需要查询所有叶子节点
B只需查询根节点
C需要合并多个节点的值,最多访问 4 个节点
D需要访问全部 7 个内部节点
3判断题

线段树的区间查询可以用于计算区间最大值、最小值、和、乘积等可结合的操作。

4填空题
以下函数实现了线段树的区间求和查询。请补充完整。

int query(int node, int l, int r, int ql, int qr) {
    if (ql <= l && r <= qr) {
        return tree[node];
    }
    int mid = (l + r) / 2;
    int res = 0;
    if (ql <= mid) res += query(node*2, l, mid, ql, qr);
    if (___) res += query(node*2+1, mid+1, r, ql, qr);
    return res;
}
5填空题
下面是一段区间最大值查询的递归实现,请填空。

int query_max(int node, int l, int r, int ql, int qr) {
    if (ql <= l && r <= qr) {
        return tree[node];
    }
    int mid = (l + r) / 2;
    int left_val = -INF, right_val = -INF;
    if (ql <= mid) left_val = query_max(node*2, l, mid, ql, qr);
    if (qr > mid) right_val = query_max(node*2+1, mid+1, r, ql, qr);
    return ___;
}