区间更新(Lazy Propagation)
前面各节中的修改查询都只影响数组中的一个元素。
但是线段树还可以对整个连续区间执行修改,同时仍然能够在
时间内完成查询。
区间加法
首先考虑最简单的问题形式:
修改查询需要给区间
中的所有数加上一个数 。
另一个需要回答的查询非常简单:
查询某个 的值。
为了让区间加法操作高效,我们在线段树每个节点中保存:
需要给该节点所对应区间中的所有数加多少。
例如,如果出现查询:
给整个数组 中的所有元素加 。
那么只需要在树的根节点中保存数字 。
通常,需要把这个数字保存到多个节点中,而这些节点所对应的区间共同组成查询区间的一个划分。
因此不需要修改全部 个值,而只需要修改
个节点。
如果之后出现查询,要求得到某个数组元素当前的值,那么只需要从根节点向该叶节点走下去,并把沿途遇到的所有值相加即可。
void build(int a[], int v, int tl, int tr) {
if (tr - tl == 1) {
t[v] = a[tl];
} else {
int tm = (tl + tr) / 2;
build(a, v*2, tl, tm);
build(a, v*2+1, tm, tr);
t[v] = 0;
}
}
void update(int v, int tl, int tr, int l, int r, int add) {
if (r <= tl || tr <= l)
return;
if (l <= tl && tr <= r) {
t[v] += add;
} else {
int tm = (tl + tr) / 2;
update(v*2, tl, tm, l, r, add);
update(v*2+1, tm, tr, l, r, add);
}
}
int get(int v, int tl, int tr, int pos) {
if (tr - tl == 1)
return t[v];
int tm = (tl + tr) / 2;
if (pos < tm)
return t[v] + get(v*2, tl, tm, pos);
else
return t[v] + get(v*2+1, tm, tr, pos);
}
区间赋值
现在假设修改操作要求把某个区间
中的每一个元素都赋值为 。
第二种查询仍然是读取数组某个元素 的值。
为了在整个区间上执行这种修改,需要在线段树的每个节点中保存:
该节点对应的区间是否全部被同一个值覆盖。
这样就可以进行一次“懒惰”的更新。
也就是说,不需要立即修改树中所有覆盖查询区间的节点,只修改其中一部分,而暂时保留其他部分不变。
如果一个节点被标记,就表示:
该节点所对应区间中的每一个元素都被赋成这个值。
事实上,它的整个子树也都应该只包含这个值。
从某种意义上来说,我们是在“偷懒”,暂时不把新值写入所有这些节点。
如果之后真的有必要,再完成这些繁琐的工作。
因此,执行修改查询后,树的某些部分会暂时变得无关紧要——其中的一些修改尚未真正向下执行。
例如,如果执行:
把整个数组 赋成某个数。
在线段树中实际上只需要修改一次:
- 把这个数写入根节点;
- 标记根节点。
其他区间仍然保持不变,尽管从逻辑上来说,这个数应该被写入整棵树。
现在假设第二次修改查询要求:
把数组前一半
赋成另一个数。
为了处理这个查询,需要给根节点的整个左子节点所对应的区间赋上这个新值。
但是在此之前,必须首先处理根节点中的信息。
这里微妙的一点在于:
数组的右半部分仍然应该保持第一次查询赋予的值,但目前右半部分的节点中并没有保存这个信息。
解决方法是把根节点的信息**向下推送(push)**到它的子节点。
也就是说,如果根节点被赋过某个值,那么:
- 把这个值赋给左子节点;
- 把这个值赋给右子节点;
- 移除根节点的标记。
完成之后,就可以给左子节点赋新的值,同时不会丢失任何必要的信息。
总结来说:
对于任何查询——无论是修改查询还是读取查询——在沿树向下移动时,都应该始终把当前节点的信息推送给两个子节点。
也可以这样理解:
在向下遍历树的过程中,我们才真正执行之前延迟的修改,而且只执行必要的部分,从而避免破坏
的复杂度。
实现中,需要编写一个 push 函数。
这个函数接收当前节点,然后把当前节点中的信息传递给它的两个子节点。
我们会在查询函数中向下递归之前调用这个函数。
不过,不需要在叶节点上调用,因为叶节点没有更下面的节点可以继续传递。
void push(int v) {
if (marked[v]) {
t[v*2] = t[v*2+1] = t[v];
marked[v*2] = marked[v*2+1] = true;
marked[v] = false;
}
}
void update(int v, int tl, int tr, int l, int r, int new_val) {
if (r <= tl || tr <= l)
return;
if (l <= tl && tr <= r) {
t[v] = new_val;
marked[v] = true;
} else {
push(v);
int tm = (tl + tr) / 2;
update(v*2, tl, tm, l, r, new_val);
update(v*2+1, tm, tr, l, r, new_val);
}
}
int get(int v, int tl, int tr, int pos) {
if (tr - tl == 1) {
return t[v];
}
push(v);
int tm = (tl + tr) / 2;
if (pos < tm)
return get(v*2, tl, tm, pos);
else
return get(v*2+1, tm, tr, pos);
}
注意:
get 函数也可以用另一种方式实现:
不真正执行延迟更新。
如果
marked[v]
为 true,那么直接返回
t[v]
即可。
区间加法,区间最大值查询
现在修改查询变成:
给某个区间中的所有元素加上一个数。
读取查询则是:
求某个区间中的最大值。
因此,对于线段树中的每个节点,需要保存其对应子区间的最大值。
比较有趣的地方在于:在执行修改操作时,如何重新计算这些值。
为此,我们为每个节点额外保存一个值。
它表示还没有向子节点传播的加数。
在进入某个子节点之前,调用 push,把这个值传递给两个子节点。
在 update 函数和 query 函数中都需要这样做。
void build(int a[], int v, int tl, int tr) {
if (tr - tl == 1) {
t[v] = a[tl];
} else {
int tm = (tl + tr) / 2;
build(a, v*2, tl, tm);
build(a, v*2+1, tm, tr);
t[v] = max(t[v*2], t[v*2 + 1]);
}
}
void push(int v) {
t[v*2] += lazy[v];
lazy[v*2] += lazy[v];
t[v*2+1] += lazy[v];
lazy[v*2+1] += lazy[v];
lazy[v] = 0;
}
void update(int v, int tl, int tr, int l, int r, int addend) {
if (r <= tl || tr <= l)
return;
if (l <= tl && tr <= r) {
t[v] += addend;
lazy[v] += addend;
} else {
push(v);
int tm = (tl + tr) / 2;
update(v*2, tl, tm, l, r, addend);
update(v*2+1, tm, tr, l, r, addend);
t[v] = max(t[v*2], t[v*2+1]);
}
}
int query(int v, int tl, int tr, int l, int r) {
if (r <= tl || tr <= l)
return -INF;
if (l <= tl && tr <= r)
return t[v];
push(v);
int tm = (tl + tr) / 2;
return max(query(v*2, tl, tm, l, r),
query(v*2+1, tm, tr, l, r));
}