CC++ & Algorithm

树状数组的区间修改与区间查询

极难2
语言版本:通用
概述:利用差分思想将区间修改问题转化为两个树状数组上的单点更新和前缀和查询,实现 O(log n) 的区间加法和区间求和。

从班级成绩看树状数组的区间修改与区间查询

想象一下,你是班长,老师让你记录全班 10 个同学的考试成绩。
突然老师说:“第 3 名到第 7 名同学每人加 5 分!”——这就是 区间修改(把连续一段都加上同一个数)。
之后老师又问:“第 2 名到第 10 名的总分是多少?”——这就是 区间查询

如果你用最笨的方法,每次区间修改都要把区间里每个人的分数挨个改一遍,查询时再把区间里的分数一个个加起来,当人数很多、操作很多时就会非常慢。
我们需要一个更聪明的方法,让每次修改和查询都只花 O(log n) 的时间 —— 这就是 树状数组的区间修改与区间查询


1. 生活中的差分思想

先回想一个常见场景:你有一本零花钱账本,每笔收入只记“这个月比上个月多了多少”,而不是记每个人具体的总数。比如:

  • 第 1 个月:妈妈给了 10 元,你记 +10
  • 第 2 个月:妈妈又给了 5 元,你记 +5
  • 第 3 个月:你没收到零花钱,你记 0(或者不记)
  • 第 4 个月:你帮家里干活得了 20 元,你记 +20

这样,到第 4 个月你想知道一共攒了多少钱,就把 1~4 月的记的数加起来:10+5+0+20 = 35 元。
这种记法就是 差分:只记录变化量,不记录绝对值。

在程序中,对原数组 a 构造一个 差分数组 d

  • d[1] = a[1]
  • d[i] = a[i] - a[i-1](i ≥ 2)

那么,给原数组的区间 [l, r] 加上 k 就等价于:

  • d[l] += k(从 l 开始多加了 k)
  • d[r+1] -= k(到 r+1 处恢复,后面不受影响)

同时,求原数组某个位置 a[i] 就等于 d[1] + d[2] + ... + d[i](差分数组的前缀和)。

比如原数组 a = [0,1,2,3,4,5](下标从1开始),差分数组就是 d = [1,1,1,1,1]
如果在区间 [2,4] 加 10,那么:

  • d[2] += 10 → 变成 11
  • d[5] -= 10 → 变成 -9

现在 d = [1,11,1,1,-9],你试试求 a[2] 应该是多少?
a[2]=d[1]+d[2]=1+11=12,而原数组在操作前 a[2]=2,加上10后确实是12,正确!
a[5]=sum(d[1..5])=1+11+1+1-9=5,原来 a[5]=5,没被影响,正确!


2. 为什么需要两个树状数组?

差分数组能轻松实现区间修改,但我们要的是 区间求和(求 a[l] 到 a[r] 的总和)。
如果只用一棵树状数组维护差分数组 d,我们能很快得到单个 a[i],但求区间和需要把区间里每个 a[i] 都加起来,复杂度还是 O(n) 。
所以我们得想一个公式,用差分数组的前缀和直接算出原数组的前缀和,从而 O(log n) 得到区间和。

公式的推导(别怕,小学数学)

我们想求 S(x) = a[1] + a[2] + ... + a[x]

因为 a[i] = d[1] + d[2] + ... + d[i](差分前缀和),所以:

S(x) = d[1] 
     + (d[1] + d[2]) 
     + (d[1] + d[2] + d[3]) 
     + ... 
     + (d[1] + d[2] + ... + d[x])

数一数每个 d[i] 出现了几次:

  • d[1] 出现在每一项,共 x 次
  • d[2] 出现在后 x-1 项,共 x-1 次
  • ...
  • d[x] 只出现在最后一项,共 1 次

所以:

S(x) = x * d[1] + (x-1) * d[2] + ... + 1 * d[x]

这个式子不好直接用树状数组求,因为系数不是常数。我们把它变形:

S(x) = (x+1) * (d[1] + d[2] + ... + d[x]) - (1*d[1] + 2*d[2] + ... + x*d[x])

也就是:

S(x) = (x+1) * sum(d[1..x]) - sum(i * d[i] for i=1..x)

解释

  • 第一个部分 (x+1) * sum(d[1..x]) 是把所有系数都变成 (x+1),多算了,所以第二部分 sum(i*d[i]) 把它减回去。
    举个例子:x=3 时,原来系数是 3,2,1(x+1)=4,所以 4,4,4 比原来多出了 (4-3)*d[1] + (4-2)*d[2] + (4-1)*d[3] = 1*d[1]+2*d[2]+3*d[3],正是第二部分!
    这个公式很巧妙,它把需要变系数的问题变成了两个固定系数的前缀和。

因此,我们只需要 两棵树状数组

  • bit1:维护差分数组 d[i] 的前缀和
  • bit2:维护 i * d[i] 的前缀和

3. 区间修改和区间查询怎么做?

区间修改(给 [l, r] 加上 k)

根据前面的差分操作:

  • bit1 上:add(l, k)add(r+1, -k)
  • bit2 上:add(l, l*k)add(r+1, -(r+1)*k)

为什么 bit2 要加 l*k-(r+1)*k
因为 i*d[i] 中,当我们给 d[l] 增加 k 时,l*d[l] 就增加了 l*k;同样,给 d[r+1] 减少 k 时,(r+1)*d[r+1] 减少了 (r+1)*k
注意:修改的是 bit2 对应的 i*d[i] 值。

区间查询(求 [l, r] 的和)

先定义函数 prefixSum(x)a[1]a[x] 的和:

prefixSum(x) = (x+1) * bit1.sum(x) - bit2.sum(x)

其中 bit1.sum(x) 返回 d[1..x] 的和,bit2.sum(x) 返回 (1*d[1] + 2*d[2] + ... + x*d[x]) 的和。
注意:如果 x=0,直接返回 0。

然后:

区间和 [l, r] = prefixSum(r) - prefixSum(l-1)

举例手算验证

还是用前面的例子:初始数组全0(长度10)。
操作1:区间 [1,5] 加1
操作2:区间 [3,7] 加2
操作3:区间 [2,4] 减1

手动算最终每个位置的值:

  • 位置1:只受操作1影响 → +1 = 1
  • 位置2:操作1 +1,操作3 -1 → 0
  • 位置3:操作1 +1,操作2 +2,操作3 -1 → 2
  • 位置4:操作1 +1,操作2 +2,操作3 -1 → 2
  • 位置5:操作1 +1,操作2 +2 → 3
  • 位置6:操作2 +2 → 2
  • 位置7:操作2 +2 → 2
  • 位置8~10:无影响 → 0

前缀和[1..3] = 1+0+2 = 3
区间和[2..6] = 0+2+2+3+2 = 9

用两棵树状数组运行程序,结果一致。


4. 代码实现(带中文注释)

C++ 版本

#include <iostream>
#include <vector>
using namespace std;

class FenwickRange {
private:
    int n;
    vector<long long> bit1; // 维护差分数组 d[i] 的前缀和
    vector<long long> bit2; // 维护 i * d[i] 的前缀和

    // lowbit 运算,树状数组基础
    long long lowbit(long long x) { return x & -x; }

    // 单点更新某个树状数组(bit)
    void add(vector<long long>& bit, int idx, long long delta) {
        while (idx <= n) {
            bit[idx] += delta;
            idx += lowbit(idx);
        }
    }

    // 查询某个树状数组的前缀和(1..idx)
    long long sum(const vector<long long>& bit, int idx) {
        long long res = 0;
        while (idx > 0) {
            res += bit[idx];
            idx -= lowbit(idx);
        }
        return res;
    }

public:
    // 构造函数:初始化大小为 n,下标从1开始
    FenwickRange(int n) : n(n), bit1(n + 1, 0), bit2(n + 1, 0) {}

    // 区间修改:给区间 [l, r] 加上 k
    void rangeAdd(int l, int r, long long k) {
        add(bit1, l, k);          // d[l] += k
        add(bit1, r + 1, -k);     // d[r+1] -= k
        add(bit2, l, l * k);      // l*d[l] 增加 l*k
        add(bit2, r + 1, -(r + 1) * k); // (r+1)*d[r+1] 减少 (r+1)*k
    }

    // 求原数组的前缀和 a[1] + ... + a[x]
    long long prefixSum(int x) {
        if (x <= 0) return 0;
        return (x + 1) * sum(bit1, x) - sum(bit2, x);
    }

    // 区间查询:求 a[l] + ... + a[r]
    long long rangeSum(int l, int r) {
        return prefixSum(r) - prefixSum(l - 1);
    }
};

int main() {
    FenwickRange ft(10); // 创建管理下标1..10的树状数组
    // 初始所有值均为0
    ft.rangeAdd(1, 5, 1);     // 第1~5位加1
    ft.rangeAdd(3, 7, 2);     // 第3~7位加2
    ft.rangeAdd(2, 4, -1);    // 第2~4位减1

    cout << "前缀和[1..3]: " << ft.prefixSum(3) << endl;   // 应输出3
    cout << "区间和[2..6]: " << ft.rangeSum(2, 6) << endl; // 应输出9

    // 手动验证各位置最终值
    // 位置: 1  2  3  4  5  6  7  8 9 10
    // 值:   1  0  2  2  3  2  2  0 0  0
    cout << "位置4的值: " << ft.rangeSum(4, 4) << endl; // 应输出2
    return 0;
}

Python 版本

class FenwickRange:
    def __init__(self, n: int):
        self.n = n
        self.bit1 = [0] * (n + 1)  # 维护差分数组 d[i] 的前缀和
        self.bit2 = [0] * (n + 1)  # 维护 i * d[i] 的前缀和

    def lowbit(self, x: int) -> int:
        return x & -x

    def _add(self, bit: list, idx: int, delta: int) -> None:
        """单点更新某个树状数组"""
        while idx <= self.n:
            bit[idx] += delta
            idx += self.lowbit(idx)

    def _sum(self, bit: list, idx: int) -> int:
        """查询某个树状数组的前缀和"""
        res = 0
        while idx > 0:
            res += bit[idx]
            idx -= self.lowbit(idx)
        return res

    def range_add(self, l: int, r: int, k: int) -> None:
        """区间修改:[l, r] 加上 k"""
        self._add(self.bit1, l, k)          # d[l] += k
        self._add(self.bit1, r + 1, -k)     # d[r+1] -= k
        self._add(self.bit2, l, l * k)      # l*d[l] 增加 l*k
        self._add(self.bit2, r + 1, -(r + 1) * k)  # (r+1)*d[r+1] 减少 (r+1)*k

    def prefix_sum(self, x: int) -> int:
        """求原数组的前缀和 a[1] + ... + a[x]"""
        if x <= 0:
            return 0
        return (x + 1) * self._sum(self.bit1, x) - self._sum(self.bit2, x)

    def range_sum(self, l: int, r: int) -> int:
        """区间查询:[l, r] 的和"""
        return self.prefix_sum(r) - self.prefix_sum(l - 1)


if __name__ == "__main__":
    ft = FenwickRange(10)  # 下标1..10,初始全0
    ft.range_add(1, 5, 1)     # 第1~5位加1
    ft.range_add(3, 7, 2)     # 第3~7位加2
    ft.range_add(2, 4, -1)    # 第2~4位减1

    print("前缀和[1..3]:", ft.prefix_sum(3))   # 输出 3
    print("区间和[2..6]:", ft.range_sum(2, 6))  # 输出 9

5. 新手容易犯的错误

  1. 下标从0还是1?
    树状数组通常用下标1开始,方便处理 lowbit。如果你用下标0,需要小心边界,建议统一用1。

  2. 忘记 r+1 可能越界
    如果 r == n,那么 r+1 = n+1 超出了树状数组大小。解决办法:在构造函数里分配 n+2 的空间,或者更新时加判断。上面的代码中,bit 大小是 n+1,而 add 的循环条件是 idx <= n,所以当 idx = n+1 时不会进入循环,但实际上我们需要更新 n+1 这个位置吗?
    注意:差分更新中,d[r+1] 是需要存在的,即使下标为 n+1 也属于辅助位置。所以树状数组的大小应该为 n+2(下标从1到n+1)。上面代码中我们把大小设为 n+1 会导致 r = nr+1 = n+1 越界(超出 n)。
    正确做法:在构造函数中分配 n + 2 的空间,使下标可用到 n+1
    例如:FenwickRange(int n) : n(n), bit1(n + 2, 0), bit2(n + 2, 0) {}
    但是,如果 r 最大为 n,那么 r+1 = n+1 是合法的,而 lowbit 操作不会超过 n+1,所以分配 n+2 就安全了。
    上面示例中为了简洁用了 n+1,实际使用时要注意。

  3. 数据范围与溢出
    公式中有 l * k(r+1) * k,如果 k 很大且 l 也大,乘法可能超过 int 范围。一定要用 long long (Python 整数无此问题)。

  4. 误把区间修改当作赋值
    这个方法只能做“加一个数”(add),不能直接“赋成某个值”(assign)。如果需要赋值,可以用线段树或差分加前缀和但更复杂。

  5. 忘记 prefixSum(l-1)l-1 可能为0
    我们在 prefixSum 里已经处理 x<=0 返回0,但小心调用时 l=1l-1=0,没问题。


6. 总结与下一步

要点说明
核心思想用差分把区间修改变成两个单点更新
两个树状数组一个维护 d[i],一个维护 i*d[i]
公式S(x) = (x+1)*sum(d[1..x]) - sum(i*d[i])
时间复杂度每次操作 O(log n)
适用场景区间加减 + 区间求和(不可赋值)

你已经学会了树状数组的区间修改与区间查询。下一步可以学习:

  • 线段树:更灵活,支持区间赋值、区间最值等
  • 矩阵中的树状数组:二维区间修改与查询
  • 树状数组 + 差分 的其他应用:如求逆序对、维护颜色数等

试着用这个方法解决一道竞赛题吧,比如“校门外的树”(区间加 1 和区间查总数),或者“借教室”(区间减某个数并判断是否够用)。加油!

例题精讲

1单选题

使用树状数组实现区间修改与区间查询时,通常需要维护几个树状数组?

A1个
B2个
C3个
D4个
2判断题

利用差分思想,只需维护一个差分树状数组即可同时实现区间加法和区间求和。

3填空题
以下是用树状数组实现区间加法和区间求和的代码片段(仅展示两个关键函数)。请填补空缺处的代码。

int bit1[N], bit2[N]; // 两个树状数组
void add(int bit[], int i, int val) { while(i <= n) { bit[i] += val; i += i & -i; } }
int sum(int bit[], int i) { int res = 0; while(i > 0) { res += bit[i]; i -= i & -i; } return res; }

void range_add(int l, int r, int val) {
    add(bit1, l, val);
    add(bit1, r+1, -val);
    ___(1)___;
    add(bit2, r+1, -val * r);
}

int prefix_sum(int i) {
    return sum(bit1, i) * i - sum(bit2, i);
}

int range_sum(int l, int r) {
    return prefix_sum(r) - prefix_sum(l-1);
}
4单选题

设两个树状数组分别为bit1(维护d[i])和bit2(维护d[i]*(i-1)),则区间[L, R]的和S(L,R)的正确计算公式是?

A(R+1)*sum(bit1,R) - sum(bit2,R) - L*sum(bit1,L-1) + sum(bit2,L-1)
BR*sum(bit1,R) - sum(bit2,R) - (L-1)*sum(bit1,L-1) + sum(bit2,L-1)
Csum(bit1,R) - sum(bit1,L-1)
Dsum(bit2,R) - sum(bit2,L-1)
5判断题

使用两个树状数组实现区间修改与区间查询时,每次区间加法和区间查询的时间复杂度都是O(log n)。