线段树

线段树常用来高效进行区间修改和区间查询,其实际上为一个二叉树,每一个节点都存储着一段区间的信息。假设 $i$ 号点维护的区间是 $[l,r]$,则其左儿子 $2i$ 维护的区间为 $[l,\left \lfloor \frac{l+r}{2} \right \rfloor ]$,其右儿子 $2i+1$ 维护的区间为 $(\left \lfloor \frac{l+r}{2} \right \rfloor,r]$。如图所示。

至于每个节点储存什么信息,就要视题目而定。

建树:

考虑递归建树。

如果递归到叶子节点,即 $l=r$,就把初始序列的对应值赋给 $tree_i$。

对于非叶子节点,递归处理其左右儿子,然后将左右儿子的信息合并起来赋值给当前节点。如何合并需要视情况而定。

void build(int s,int l,int r){//s为当前节点,l,r为当前节点所维护的区间
	if(l==r){
		tree[s]=a[l];
		return;
	}
	int mid=(l+r)>>1;
	build(s*2,l,mid);//递归左儿子
	build(s*2+1,mid+1,r);//递归右儿子
	tree[s]=merge(tree[s*2],tree[s*2+1]);//合并左右儿子信息
}

合并:

合并过程就是用两个儿子所维护的区间信息推出当前节点所维护的区间信息。

比如维护区间最大值,则当前节点维护的区间最大值就是两个儿子所维护的区间最大值中的最大值。

int merge(int l,int r){
    return max(l,r);
}

更新:

如果要更新 $a_x$ 的信息,我们就要从区间 $[1,n]$ 递归下去,每次判断 $x$ 在左儿子内还是右儿子内,然后递归对应的区间。当 $l=r$ 时,说明当前节点就是 $x$,更新节点信息,然后返回,一路将信息更新回去。

以下是实现将 $a_{pos}$ 更新为 $val$ 的代码。

void update(int s,int l,int r,int pos,int val){
	if(l==r){
		tree[s]=val;
		return;
	}
	int mid=(l+r)>>1;
	if(mid>=pos){
		update(s*2,l,mid,pos,val);
	}
	else update(s*2+1,mid+1,r,pos,val);
	tree[s]=merge(tree[s*2],tree[s*2+1]);
}

查询:

查询的思路其实跟更新差不多,都是不断判断目标区间在左儿子还是右儿子内,然后递归那个儿子,当当前区间完全包含在目标区间中时,直接返回当前节点维护的值即可。如果目标区间横跨左右儿子,则需要递归左右儿子,然后将左右儿子的信息合并起来。

int query(int s,int l,int r,int ql,int qr){
	if(l>=ql&&r<=qr){
		return tree[s];
	}
	int mid=(l+r)>>1;
	if(qr<=mid){
		return query(s*2,l,mid,ql,qr);
	}
	if(ql>mid){
		return query(s*2+1,mid+1,r,ql,qr);
	}
	int ll=query(s*2,l,mid,ql,qr);
	int rr=query(s*2+1,mid+1,r,ql,qr);
	return merge(ll,rr);
}

当然,以上的内容都只能实现单点修改。

简单的线段树。

例题:小白逛公园

让我们维护区间内最大子段和。每个节点只需要维护当前区间的最大前缀和、最大后缀和、区间和以及最大子段和,用这些信息即可更新父亲节点的值。

#include<bits/stdc++.h>
using namespace std;
int n,q;
int a[500005];
struct node{
	int maxn,sum,pre,suf;
}tree[800005];
node merge(node l,node r){
	node res;
	res.sum=l.sum+r.sum;
	res.pre=max(l.pre,l.sum+r.pre);
	res.suf=max(r.suf,r.sum+l.suf);
	res.maxn=max({l.maxn,r.maxn,l.suf+r.pre});
	return res;
}
void build(int s,int l,int r){
	if(l==r){
		tree[s]={a[l],a[l],a[l],a[l]};
		return;
	}
	int mid=(l+r)>>1;
	build(s*2,l,mid);
	build(s*2+1,mid+1,r);
	tree[s]=merge(tree[s*2],tree[s*2+1]);
}
void update(int s,int l,int r,int pos,int val){
	if(l==r){
		tree[s]={val,val,val,val};
		return;
	}
	int mid=(l+r)>>1;
	if(mid>=pos){
		update(s*2,l,mid,pos,val);
	}
	else update(s*2+1,mid+1,r,pos,val);
	tree[s]=merge(tree[s*2],tree[s*2+1]);
}
node query(int s,int l,int r,int ql,int qr){
	if(ql<=l&&qr>=r){
		return tree[s];
	}
	int mid=(l+r)>>1;
	if(qr<=mid){
		return query(s*2,l,mid,ql,qr);
	}
	if(ql>mid){
		return query(s*2+1,mid+1,r,ql,qr);
	}
	node ll=query(s*2,l,mid,ql,qr);
	node rr=query(s*2+1,mid+1,r,ql,qr);
	return merge(ll,rr);
}
signed main(){
	ios::sync_with_stdio(0);
	cin.tie(0),cout.tie(0);
	cin>>n>>q;
	for(int i=1;i<=n;i++){
		cin>>a[i];
	}
	build(1,1,n);
	while(q--){
		int op,x,d,l,r;
		cin>>op;
		if(!(op&1)){
			cin>>x>>d;
			update(1,1,n,x,d);
		}
		else{
			cin>>l>>r;
			if(l>r)swap(l,r);
			cout<<query(1,1,n,l,r).maxn<<'\n';
		}
	}
	return 0;
}

懒标记

如果我们需要进行区间修改,那么前面的更新方式就太慢了。

此时,我们可以使用懒标记。懒标记的核心思想就是需要当前节点时再更新,如果不用当前节点就暂时不更新。这样的复杂度是 $O(logN)$ 的。

一个节点的懒标记就代表这个节点所维护的区间要进行的变化量。

在进行更新或查询操作时,需要下推懒标记,即根据懒标记对当前区间进行更新。

这是维护区间最值,且支持区间修改的代码:

#include<bits/stdc++.h>
#define int long long
using namespace std;
int n,q;
int a[200005];
struct node{
	int maxn,lazy_tag;
}tree[800005];
void build(int s,int l,int r){
	if(l==r){
		tree[s].maxn=a[l];
		return;
	}
	int mid=(l+r)>>1;
	build(s*2,l,mid);
	build(s*2+1,mid+1,r);
	tree[s].maxn=max(tree[s*2].maxn,tree[s*2+1].maxn);
}
void lazy(int s,int val){
	tree[s].maxn+=val;
	tree[s].lazy_tag+=val;
}
void push_down(int s){
	if(tree[s].lazy_tag!=0){
		lazy(s*2,tree[s].lazy_tag);
		lazy(s*2+1,tree[s].lazy_tag);
		tree[s].lazy_tag=0;
	}
}
void update(int s,int l,int r,int ql,int qr,int val){
	if(l>=ql&&r<=qr){
		lazy(s,val);
		return;
	}
	push_down(s);
	int mid=(l+r)>>1;
	if(ql<=mid){
		update(s*2,l,mid,ql,qr,val);
	}
	if(qr>mid){
		update(s*2+1,mid+1,r,ql,qr,val);
	}
	tree[s].maxn=max(tree[s*2].maxn,tree[s*2+1].maxn);
}
int query(int s,int l,int r,int ql,int qr){
	if(l>=ql&&r<=qr){
		return tree[s].maxn;
	}
	push_down(s);
	int mid=(l+r)>>1;
	if(qr<=mid){
		return query(s*2,l,mid,ql,qr);
	}
	if(ql>mid){
		return query(s*2+1,mid+1,r,ql,qr);
	}
	int ll=query(s*2,l,mid,ql,qr);
	int rr=query(s*2+1,mid+1,r,ql,qr);
	return max(ll,rr);
}
signed main(){
	cin>>n>>q;
	for(int i=1;i<=n;i++){
		cin>>a[i];
	}
	build(1,1,n);
	while(q--){
		int op,l,r,d;
		cin>>op;
		if(op&1){
			cin>>l>>r>>d;
			update(1,1,n,l,r,d);
		}
		else{
			cin>>l>>r;
			cout<<query(1,1,n,l,r)<<'\n';
		}
	}
	return 0;
}