ACM 学习篇 11:线段树(Segment Tree)

上一篇讲了树状数组,它能高效处理单点修改和区间求和。但如果问题需要:

  • 区间修改:把一个区间内的所有元素都加上某个值;
  • 区间最值查询:查询一个区间的最大值或最小值;
  • 更复杂的区间操作:区间覆盖、区间翻转等。

树状数组就不够用了。这时候需要 线段树(Segment Tree)——一种更通用、更强大的数据结构。

线段树的核心思想

线段树是一种二叉树,每个节点代表一个区间 [l, r]

  • 叶子节点l = r,存储单个元素的值;
  • 非叶子节点:存储其左右子树区间的合并信息(如和、最值等);
  • 结构:通常用数组存储,假设原数组长度为 n,线段树数组大小一般取 4 * n

比如数组 [1, 3, 5, 7, 9, 11] 的线段树结构:

                [1-6] sum=36
              /         \
        [1-3] sum=9    [4-6] sum=27
        /    \          /    \
    [1-2]   [3]    [4-5]   [6]
    /   \          /   \
  [1]  [2]      [4]   [5]
   1    3        7     9

线段树的基本操作

1. 单点更新

a[i] 改成 val,需要从叶子节点一路更新到根节点。

void update(int node, int l, int r, int idx, int val) {
    if (l == r) {
        tree[node] = val;
        return;
    }
    int mid = (l + r) / 2;
    if (idx <= mid) {
        update(2 * node, l, mid, idx, val);
    } else {
        update(2 * node + 1, mid + 1, r, idx, val);
    }
    // 合并左右子树的信息
    tree[node] = tree[2 * node] + tree[2 * node + 1];
}

2. 区间查询

查询区间 [ql, qr] 的和,分三种情况:

  • 当前节点区间完全在查询区间外:返回 0(或其他单位元);
  • 当前节点区间完全在查询区间内:直接返回该节点的值;
  • 当前节点区间与查询区间部分重叠:递归查询左右子树并合并结果。
int query(int node, int l, int r, int ql, int qr) {
    if (qr < l || ql > r) return 0;  // 完全不重叠
    if (ql <= l && r <= qr) return tree[node];  // 完全包含
    int mid = (l + r) / 2;
    return query(2 * node, l, mid, ql, qr) + query(2 * node + 1, mid + 1, r, ql, qr);
}

3. 区间修改与懒标记

如果需要把区间 [ul, ur] 的所有元素都加上 delta,朴素做法需要遍历每个叶子节点,复杂度是 O(n)。这太慢了。

懒标记(Lazy Propagation) 的思想是:先把修改记录在当前节点的懒标记中,等到需要访问子节点时再把标记传递下去。这样可以把区间修改的复杂度降到 O(log n)。

void pushDown(int node, int l, int r) {
    if (lazy[node] != 0 && l != r) {
        // 把懒标记传递给左右子节点
        int mid = (l + r) / 2;
        tree[2 * node] += lazy[node] * (mid - l + 1);
        lazy[2 * node] += lazy[node];
        tree[2 * node + 1] += lazy[node] * (r - mid);
        lazy[2 * node + 1] += lazy[node];
        // 清除当前节点的懒标记
        lazy[node] = 0;
    }
}

void rangeUpdate(int node, int l, int r, int ul, int ur, int delta) {
    if (ur < l || ul > r) return;
    if (ul <= l && r <= ur) {
        tree[node] += delta * (r - l + 1);
        lazy[node] += delta;
        return;
    }
    pushDown(node, l, r);  // 先传递懒标记
    int mid = (l + r) / 2;
    rangeUpdate(2 * node, l, mid, ul, ur, delta);
    rangeUpdate(2 * node + 1, mid + 1, r, ul, ur, delta);
    tree[node] = tree[2 * node] + tree[2 * node + 1];
}

4. 区间修改后的查询

查询时也要先传递懒标记,确保访问的子节点信息是最新的:

int queryWithLazy(int node, int l, int r, int ql, int qr) {
    if (qr < l || ql > r) return 0;
    if (ql <= l && r <= qr) return tree[node];
    pushDown(node, l, r);  // 查询前先传递懒标记
    int mid = (l + r) / 2;
    return queryWithLazy(2 * node, l, mid, ql, qr) + queryWithLazy(2 * node + 1, mid + 1, r, ql, qr);
}

完整模板:区间求和

class SegmentTree {
private:
    vector tree;
    vector lazy;
    int n;

    void build(int node, int l, int r, vector& a) {
        if (l == r) {
            tree[node] = a[l];
            return;
        }
        int mid = (l + r) / 2;
        build(2 * node, l, mid, a);
        build(2 * node + 1, mid + 1, r, a);
        tree[node] = tree[2 * node] + tree[2 * node + 1];
    }

    void pushDown(int node, int l, int r) {
        if (lazy[node] != 0 && l != r) {
            int mid = (l + r) / 2;
            tree[2 * node] += lazy[node] * (mid - l + 1);
            lazy[2 * node] += lazy[node];
            tree[2 * node + 1] += lazy[node] * (r - mid);
            lazy[2 * node + 1] += lazy[node];
            lazy[node] = 0;
        }
    }

    void update(int node, int l, int r, int idx, int val) {
        if (l == r) {
            tree[node] = val;
            return;
        }
        pushDown(node, l, r);
        int mid = (l + r) / 2;
        if (idx <= mid) {
            update(2 * node, l, mid, idx, val);
        } else {
            update(2 * node + 1, mid + 1, r, idx, val);
        }
        tree[node] = tree[2 * node] + tree[2 * node + 1];
    }

    void rangeUpdate(int node, int l, int r, int ul, int ur, int delta) {
        if (ur < l || ul > r) return;
        if (ul <= l && r <= ur) {
            tree[node] += delta * (r - l + 1);
            lazy[node] += delta;
            return;
        }
        pushDown(node, l, r);
        int mid = (l + r) / 2;
        rangeUpdate(2 * node, l, mid, ul, ur, delta);
        rangeUpdate(2 * node + 1, mid + 1, r, ul, ur, delta);
        tree[node] = tree[2 * node] + tree[2 * node + 1];
    }

    int query(int node, int l, int r, int ql, int qr) {
        if (qr < l || ql > r) return 0;
        if (ql <= l && r <= qr) return tree[node];
        pushDown(node, l, r);
        int mid = (l + r) / 2;
        return query(2 * node, l, mid, ql, qr) + query(2 * node + 1, mid + 1, r, ql, qr);
    }

public:
    SegmentTree(vector& a) {
        n = a.size() - 1;  // 假设 a 从 1 开始
        tree.resize(4 * (n + 1));
        lazy.resize(4 * (n + 1), 0);
        build(1, 1, n, a);
    }

    void update(int idx, int val) {
        update(1, 1, n, idx, val);
    }

    void rangeUpdate(int l, int r, int delta) {
        rangeUpdate(1, 1, n, l, r, delta);
    }

    int query(int l, int r) {
        return query(1, 1, n, l, r);
    }
};

区间最值线段树

线段树不仅能维护区间和,也能维护区间最大值/最小值。只需要修改合并逻辑:

class SegmentTreeMax {
private:
    vector tree;
    vector lazy;
    int n;

    void build(int node, int l, int r, vector& a) {
        if (l == r) {
            tree[node] = a[l];
            return;
        }
        int mid = (l + r) / 2;
        build(2 * node, l, mid, a);
        build(2 * node + 1, mid + 1, r, a);
        tree[node] = max(tree[2 * node], tree[2 * node + 1]);
    }

    void pushDown(int node) {
        if (lazy[node] != 0) {
            tree[2 * node] += lazy[node];
            lazy[2 * node] += lazy[node];
            tree[2 * node + 1] += lazy[node];
            lazy[2 * node + 1] += lazy[node];
            lazy[node] = 0;
        }
    }

    void rangeUpdate(int node, int l, int r, int ul, int ur, int delta) {
        if (ur < l || ul > r) return;
        if (ul <= l && r <= ur) {
            tree[node] += delta;
            lazy[node] += delta;
            return;
        }
        pushDown(node);
        int mid = (l + r) / 2;
        rangeUpdate(2 * node, l, mid, ul, ur, delta);
        rangeUpdate(2 * node + 1, mid + 1, r, ul, ur, delta);
        tree[node] = max(tree[2 * node], tree[2 * node + 1]);
    }

    int query(int node, int l, int r, int ql, int qr) {
        if (qr < l || ql > r) return INT_MIN;
        if (ql <= l && r <= qr) return tree[node];
        pushDown(node);
        int mid = (l + r) / 2;
        return max(query(2 * node, l, mid, ql, qr), query(2 * node + 1, mid + 1, r, ql, qr));
    }

public:
    // ... 构造函数和接口函数类似
};

典型题:区间染色

问题:初始时所有位置都是白色,多次操作:

  • 把区间 [l, r] 染成某种颜色;
  • 查询区间 [l, r] 有多少种不同的颜色。

思路:用线段树维护区间的颜色信息,懒标记表示区间被统一染色。

class ColorSegmentTree {
private:
    vector tree;  // -1 表示区间颜色不统一
    vector lazy;
    int n;

    void pushDown(int node) {
        if (lazy[node] != -1) {
            tree[2 * node] = lazy[node];
            lazy[2 * node] = lazy[node];
            tree[2 * node + 1] = lazy[node];
            lazy[2 * node + 1] = lazy[node];
            lazy[node] = -1;
        }
    }

    void update(int node, int l, int r, int ul, int ur, int color) {
        if (ur < l || ul > r) return;
        if (ul <= l && r <= ur) {
            tree[node] = color;
            lazy[node] = color;
            return;
        }
        pushDown(node);
        int mid = (l + r) / 2;
        update(2 * node, l, mid, ul, ur, color);
        update(2 * node + 1, mid + 1, r, ul, ur, color);
        // 如果左右子树颜色相同,合并
        if (tree[2 * node] != -1 && tree[2 * node] == tree[2 * node + 1]) {
            tree[node] = tree[2 * node];
        } else {
            tree[node] = -1;
        }
    }

    void query(int node, int l, int r, int ql, int qr, unordered_set& colors) {
        if (qr < l || ql > r) return;
        if (tree[node] != -1) {
            colors.insert(tree[node]);
            return;
        }
        pushDown(node);
        int mid = (l + r) / 2;
        query(2 * node, l, mid, ql, qr, colors);
        query(2 * node + 1, mid + 1, r, ql, qr, colors);
    }

public:
    ColorSegmentTree(int size) : n(size) {
        tree.resize(4 * (n + 1), 0);  // 初始颜色为 0(白色)
        lazy.resize(4 * (n + 1), -1);
    }

    void paint(int l, int r, int color) {
        update(1, 1, n, l, r, color);
    }

    int countColors(int l, int r) {
        unordered_set colors;
        query(1, 1, n, l, r, colors);
        return colors.size();
    }
};

线段树的扩展应用

  • 二维线段树:处理二维区间查询和修改;
  • 动态开点线段树:处理值域很大但实际点数不多的情况;
  • 可持久化线段树(主席树):维护历史版本,支持区间第 K 小;
  • 线段树合并/分裂:动态维护多个线段树;
  • 扫描线 + 线段树:处理平面矩形面积并等问题。

树状数组 vs 线段树对比

特性 树状数组 线段树
单点修改 O(log n) O(log n)
区间查询(和) O(log n) O(log n)
区间修改(加) 需要差分技巧 O(log n)(懒标记)
区间最值查询 不支持 O(log n)
区间覆盖 不支持 O(log n)
代码长度 短(~20 行) 长(~50 行)
常数 较大
空间 O(n) O(4n)
灵活性 只能处理可合并的信息 可以处理更复杂的操作

线段树的易错点

  • 数组下标从 1 开始:线段树通常用 1 作为根节点,这样左子节点是 2*node,右子节点是 2*node+1;
  • 懒标记的初始化:懒标记初始值要合理(如求和用 0,最值用 -INF/INF,颜色用 -1);
  • pushDown 的时机:在更新和查询前都要调用 pushDown,确保子节点信息正确;
  • 区间合并逻辑:不同问题的合并方式不同(求和、最值、计数等);
  • 线段树大小:通常开 4*n 的空间,防止越界;
  • 递归边界:注意 l == r 时的处理,避免死循环。

什么时候用线段树

看到这些关键词,优先考虑线段树:

  • 区间修改 + 区间查询(树状数组需要差分技巧);
  • 区间最值查询(树状数组不支持);
  • 区间覆盖、翻转、赋值(需要特殊的懒标记);
  • 需要维护多个信息(如同时维护区间和、最大值、最小值);
  • 可持久化(主席树)。

这一篇先记住什么

  • 线段树结构:每个节点代表一个区间,叶子节点是单个元素;
  • 单点更新:从叶子到根,一路更新;
  • 区间查询:分三种情况(完全不重叠、完全包含、部分重叠);
  • 懒标记:延迟传递修改,把区间修改降到 O(log n);
  • pushDown:查询或更新前传递懒标记;
  • 空间:开 4*n 的数组;
  • 对比树状数组:树状数组代码短、常数小;线段树更灵活、支持更多操作。

下一篇继续数据结构:ST 表(稀疏表),它能在 O(1) 时间内回答区间最值查询,但不支持修改。