树状数组的单点修改与区间查询
较难2树状数组入门:单点修改与区间查询的魔法
想象一下,你是班上的生活委员,每天要记录每个同学的零花钱(比如小明的零花钱是 20 元,小红是 15 元……)。突然有一天,老师问:“从第 3 个同学到第 8 个同学,他们的零花钱总和是多少?” 你拿出本子,一个个加起来,刚算完,老师又说:“小强今天涨了 5 元零花钱,快更新一下!” 你赶紧涂改,然后老师又问另一个区间…… 这样反复几十次,你是不是快疯了?
别急,这时候 树状数组(也叫 Fenwick 树)就能帮上大忙。它是一个 既能快速修改单个元素,又能快速求任意连续区间和 的数据结构。每次操作只需要 O(log n) 的时间,哪怕有 1000 个同学,也只需要几十步就算出来了。
树状数组最基础的两个操作就是:
- 单点修改:修改数组中的一个元素(比如把某个同学的零花钱增加或减少)。
- 区间查询:求任意一段连续子数组的和(比如第 L 到第 R 个同学的零花钱总和)。
下面我们就从生活中的例子出发,一步步揭开它的秘密。
从生活中的例子引入
你每天记录班上同学的身高(下表从 1 开始编号):
| 同学编号 | 1 | 2 | 3 | 4 | 5 | 6 | 7 | 8 |
|---|---|---|---|---|---|---|---|---|
| 身高(cm) | 150 | 155 | 148 | 162 | 170 | 158 | 165 | 172 |
现在老师想知道第 3 个到第 8 个同学的平均身高,你需要快速算出他们的身高总和。更麻烦的是,班上同学可能会突然长高(单点修改),老师告诉你某个同学的新身高,你要把原来的记录更新。如果每天都要更新很多次,又经常查询不同区间的总和,有什么好办法呢?
你已经知道树状数组可以快速计算 前缀和(前 x 个元素的和),而 区间和 = 前缀和(R) - 前缀和(L-1)。同时,当某个同学身高变化时,我们只需要更新所有包含这个同学的“管理者”节点,也就是 add(idx, delta) 操作。这就是树状数组的单点修改与区间查询。
核心思想:两个小动作搞定一切
1. lowbit 是什么?为什么它这么神奇?
树状数组的每个节点 tree[i] 负责管理原数组中的一段连续区间。这个区间的长度正好是 lowbit(i),也就是 i 的二进制表示中 最低位 1 所代表的值。
- 例如:
i=3,二进制11,最低位 1 是1,所以lowbit(3) = 1。tree[3]只管理自己(区间[3,3])。 i=4,二进制100,最低位 1 是100(即4),所以lowbit(4) = 4。tree[4]管理区间[1,4]。i=6,二进制110,最低位 1 是10(即2),所以lowbit(6) = 2。tree[6]管理区间[5,6]。
怎么计算 lowbit? 代码里通常用 x & -x。因为负数的二进制是补码,-x 等于 ~x + 1,x & -x 就能取出最低位的 1 和后面的 0。比如 6 & -6:
6 = 0000 0110
-6 = 1111 1010(补码)
& = 0000 0010 = 2
2. 单点修改:从下往上“通知”所有管理者
如果我们要把 a[3] 增加 delta,哪些树状数组节点需要更新呢?就是所有管理区间 包含下标 3 的节点。从 i=3 开始,每次加上 lowbit(i),就能找到下一个更大的管理者。
- 当前节点
3:lowbit(3)=1,更新tree[3]。 - 下一个节点
3+1=4:lowbit(4)=4,更新tree[4]。 - 再下一个节点
4+4=8:lowbit(8)=8,更新tree[8]。 - 再下一个
8+8=16 > n=8,停止。
是不是很像爬楼梯?每次跨一个 lowbit 大小的台阶,直到顶层。
生活中的比喻:假设每个同学都有一个“小队长”负责收集自己小队的身高总和。a[3] 是队员 3,他的直接小队长是节点 3(只管自己),大队长是节点 4(管 14),总队长是节点 8(管 18)。队员 3 长高了,他要向所有管他的队长报告,队长们更新自己的记录。
3. 区间查询:从右向左“收集”小队总和
要求前 7 个同学的身高总和,即 sum(7)。我们从 i=7 开始,每次减去 lowbit(i),累加对应节点的值。
i=7,lowbit(7)=1,加上tree[7](管理[7,7])。i=6(因为7-1=6),lowbit(6)=2,加上tree[6](管理[5,6])。i=4(因为6-2=4),lowbit(4)=4,加上tree[4](管理[1,4])。i=0,停止。
我们收集到的三个区间 [7,7]、[5,6]、[1,4] 恰好不重不漏地覆盖了 [1,7]。
生活中的比喻:总队长想知道前 7 个队员的身高总和,他先问小队长 7(管自己),再问小队长 6(管 56),再问大队长 4(管 14),把所有答案加起来就行。而区间查询 [3,8] 就相当于先问总队长前 8 个的总和,再减去前 2 个的总和。
ASCII 图来帮忙
更新 a[3] 时,需要更新的节点(n=8)
tree[1] tree[2] tree[3] tree[4] tree[5] tree[6] tree[7] tree[8]
[1,1] [1,2] [3,3] [1,4] [5,5] [5,6] [7,7] [1,8]
更新: 受影响 受影响 受影响 受影响
- 节点 3 更新(因为它直接管 3 自己)
- 节点 4 更新(因为 3+1=4,它管 1~4)
- 节点 8 更新(因为 4+4=8,它管 1~8)
查询 sum(7) 时,累加的节点
sum(7)=tree[7]+tree[6]+tree[4]
7 -> tree[7] (lowbit=1) → 覆盖 [7,7]
6 -> tree[6] (7-1=6, lowbit=2) → 覆盖 [5,6]
4 -> tree[4] (6-2=4, lowbit=4) → 覆盖 [1,4]
0 -> 停止
把这些区间拼起来就是 [1,4] ∪ [5,6] ∪ [7,7] = [1,7],完美!
完整的代码实现
我们已经了解了 add 和 sum 两个核心函数,现在直接使用它们实现单点修改和区间查询。
C++ 完整代码
#include <iostream>
#include <vector>
using namespace std;
class Fenwick {
private:
int n; // 数组大小
vector<int> tree; // 树状数组,下标从1开始
int lowbit(int x) { return x & -x; } // 计算 lowbit
public:
// 构造函数,初始化大小为 n,所有值为0
Fenwick(int n) : n(n), tree(n + 1, 0) {}
// 单点修改:在下标 idx 处增加 delta
void add(int idx, int delta) {
while (idx <= n) {
tree[idx] += delta;
idx += lowbit(idx);
}
}
// 前缀和:返回前 idx 个元素的和
int sum(int idx) {
int res = 0;
while (idx > 0) {
res += tree[idx];
idx -= lowbit(idx);
}
return res;
}
// 区间查询:返回 [l, r] 的和
int rangeSum(int l, int r) {
if (l > r) return 0;
return sum(r) - sum(l - 1);
}
};
int main() {
int a[] = {0, 2, 5, 1, 3, 7, 8, 4, 6}; // 原数组,下标从1开始,a[0]不用
int n = 8;
Fenwick ft(n);
// 初始化:逐个调用 add 把原数组的值添加进去
for (int i = 1; i <= n; ++i) {
ft.add(i, a[i]);
}
cout << "初始区间[2,5]的和: " << ft.rangeSum(2,5) << endl; // 5+1+3+7=16
// 单点修改:将 a[3] 改为 10 (原值为1,增加9)
ft.add(3, 9);
cout << "修改后区间[2,5]的和: " << ft.rangeSum(2,5) << endl; // 16+9=25
// 再次修改:将 a[7] 减2
ft.add(7, -2);
cout << "再次修改后区间[2,5]的和: " << ft.rangeSum(2,5) << endl; // 25不变,因为7不在区间[2,5]
cout << "区间[6,8]的和: " << ft.rangeSum(6,8) << endl; // 原为7+8+4+6=25,减去2得23
return 0;
}
Python 完整代码
class Fenwick:
def __init__(self, n):
self.n = n
self.tree = [0] * (n + 1) # 树状数组,下标从1开始
def lowbit(self, x):
return x & -x
def add(self, idx, delta):
# 单点修改:在下标 idx 处增加 delta
while idx <= self.n:
self.tree[idx] += delta
idx += self.lowbit(idx)
def sum(self, idx):
# 前缀和:返回前 idx 个元素的和
res = 0
while idx > 0:
res += self.tree[idx]
idx -= self.lowbit(idx)
return res
def range_sum(self, l, r):
# 区间查询:返回 [l, r] 的和
if l > r:
return 0
return self.sum(r) - self.sum(l - 1)
if __name__ == "__main__":
a = [0, 2, 5, 1, 3, 7, 8, 4, 6] # 原数组,下标1-based
n = 8
ft = Fenwick(n)
# 初始化:逐个添加
for i in range(1, n+1):
ft.add(i, a[i])
print("初始区间[2,5]的和:", ft.range_sum(2,5)) # 16
ft.add(3, 9) # 将 a[3] 从1改成10,增加9
print("修改后区间[2,5]的和:", ft.range_sum(2,5)) # 25
ft.add(7, -2)
print("再次修改后区间[2,5]的和:", ft.range_sum(2,5)) # 25
print("区间[6,8]的和:", ft.range_sum(6,8)) # 23 (7+8+4+6-2)
运行上面的代码,你会得到:
初始区间[2,5]的和: 16
修改后区间[2,5]的和: 25
再次修改后区间[2,5]的和: 25
区间[6,8]的和: 23
新手最容易犯的 4 个错误
-
下标从 1 开始,却写成了 0
树状数组通常用tree[1]到tree[n],下标 0 留空。如果你用add(0, delta)或者sum(0),会陷入死循环(因为lowbit(0) = 0,永远加/减 0)。所以切记:数组下标从 1 开始。 -
区间查询时忘记减 1
想求[l, r]的和,必须用sum(r) - sum(l-1)。有人写成sum(r) - sum(l),那就少算了a[l]。例如求[2,5],sum(5)是前5个,sum(2)是前2个,差是第3、4、5个,丢掉了第2个。正确做法是减sum(l-1)。 -
单点修改时用新值代替增量
add函数期望传入的是 改变量 delta,而不是新值。比如原来a[3]=10,现在要改成8,你应该调用add(3, -2),而不是add(3, 8)。若直接传入 8,相当于新增了 8,而不是替换。 -
忘记初始化树状数组
刚建立的树状数组所有节点都是 0,需要把原数组的值通过add添加进去才能正确查询。有人直接用一个循环for i in 1..n: ft.add(i, a[i])来初始化,这是 O(n log n) 的。其实还有一种 O(n) 的初始化方法,但初学用循环就够。
总结要点
- 单点修改:调用
add(idx, delta),内部通过不断加 lowbit 更新所有受影响的节点。 - 区间查询:通过前缀和之差
sum(r) - sum(l-1)实现,sum函数通过不断减 lowbit 累加。 - 两个操作的时间复杂度都是 O(log n),比暴力法的 O(n) 快很多,适合频繁修改和查询的场景。
- 构建树状数组需要 O(n log n) 或 O(n) 时间,但之后每次操作很快。
- 注意下标从 1 开始,0 下标不使用。
- 树状数组支持的值可以是任意可加的类型(如整数、取模域),但必须是可累加的。
通过这个例子,你应该理解了树状数组最经典的应用。下一步,我们可以用这个技巧解决更复杂的问题,比如 求逆序对(后面会讲),或者实现 区间修改和单点查询(利用差分思想)。试试自己手算一下 ft.sum(5) 的过程,加深理解。
相关指引
如果你已经掌握了单点修改与区间查询,下一个要学习的可能是:
- 树状数组的区间修改与单点查询(利用差分数组)
- 树状数组的区间修改与区间查询(用两个树状数组)
- 树状数组求逆序对(统计每个数前面有多少个比它大的)
- 如果数据范围很大,还可以用 离散化 配合树状数组
继续加油,你离数据结构高手又近了一步!
例题精讲
在树状数组中进行单点修改(将原数组第i个元素增加x)时,需要更新树状数组中哪些位置?
已知树状数组t已构建完毕,要查询区间[l, r](1≤l≤r≤n)的和,正确的实现方式是?
树状数组的单点修改和区间查询操作的时间复杂度都是O(log n)。
如果只对树状数组进行单点修改和区间查询,那么原数组可以完全省略,只保留树状数组即可。
以下代码实现了树状数组的单点修改和区间查询。请在空白处填入正确的表达式。
int n;
int t[100005];
void add(int pos, int val) {
for (int i = pos; i <= n; i += ___ ) {
t[i] += val;
}
}
int sum(int pos) {
int res = 0;
for (int i = pos; i > 0; i -= ___ ) {
res += t[i];
}
return res;
}
int range_sum(int l, int r) {
return sum(r) - sum(l-1);
}