浅谈主席树
我讨厌史山数据结构!呜呜呜。
引入
给定一个长度为 \(n\) 的正整数序列 \(a\),有 \(q\) 次查询,每次查询序列 \(a\) 中下标为 \([l,r]\) 的区间中第 \(k\) 小的数字的值。
如果 \(n,q \le 10^3\),这就是一个简单的暴力题了;但当 \(n,q \le 10^5\) 时,暴力就没办法解决问题了,这个时候我们该怎么办呢?
考虑在建树的时候保存每次插入的历史版本,方便查找第 \(k\) 小值。这就是主席树的核心骨干了。
什么是主席树?
主席树,全称可持久化权值线段树,是可持久化线段树中的一种。它是一种支持历史版本查询的数据结构。假如你对一个数组做了多次修改,主席树可以让你高效地查询“第 \(k\) 次修改后的某个区间和是多少”,而不用把每次修改都存一份完整数据。
为什么叫主席树?
主席树最初是由黄嘉泰在 2010 年独立提出并推广的,由于黄嘉泰的姓名首拼 hjt 与当时中华人民共和国主席一样,于是该算法就被称为“主席树”了。
主席树怎么做?
主席树的核心骨干在刚才的引入部分已经提到过了:
考虑在建树的时候保存每次插入的历史版本,方便查找第 \(k\) 小值。
可是可是可是,你该怎么保存呢?
简单暴力一点,每个都开棵线段树——可那样空间不得炸掉??
分析一下,其实可以发现,每次修改操作修改的点的个数是一样的!只改了 \(O(\log n)\) 个节点,形成一条从根到叶子的链,也就是说每次修改的节点个数其实就是树的高度。
从 OI-wiki 偷借了一个图来:

注意主席树不能使用堆式存储法,也就是不能直接草率地用 \(2x\) 和 \(2x+1\) 来表示左右儿子,而是应该动态开点,并保存每个节点的左右儿子编号。
所以我们只要在记录左右儿子的基础上,保存插入每个数的时候的根节点就可以实现持久化啦~
但是这样好像还是不行。我们把问题简化一下,不求 \([l,r]\) 的第 \(k\) 小值了,而是去求 \([1,r]\) 的。这个非常好做,只需要找到插入 \(r\) 时的根节点版本,然后用普通权值线段树(也叫做值域线段树)就行了。
怎么找插入 \(r\) 时的根节点版本?简单,对于每个节点 \(x\),插入时维护根节点 \(Rt_x\) 即可。
回顾原题,怎么解决非前缀的第 \(k\) 小值呢?这里我们只需要用上前缀和——它的本质是巧妙运用了区间减法的性质,通过预处理从而达到 \(O(1)\) 回答每个询问。而在这里,主席树统计的信息也满足这个性质!所以,如果需要得到 \([l,r]\) 的答案,只需要用 \([1,r]\) 的答案减去 \([1,l-1]\) 的答案就可以啦 > <
最后算算空间吧。由于动态开点,一棵树最多 \(2n-1\) 个节点;\(n\) 次修改,每次至多增加 \(\log_2 n +1\) 个节点;\(n \le 10^5\) 的数据范围,最后粗略估计下来也就 \(2 \times 10^6\) 左右,完全没有问题啦!
静态查询第 \(k\) 小实现
#include<bits/stdc++.h>
#define LL long long
#define UInt unsigned int
#define ULL unsigned long long
#define LD long double
#define pii pair<int,int>
#define pLL pair<LL,LL>
#define pDD pair<LD,LD>
#define fr first
#define se second
#define pb push_back
#define isr insert
using namespace std;
const int N = 2e5+5;
const int K = (N<<5);
int n,m,Q,p[N],a[N];
int node_cnt,Rt[N],ls[K],rs[K],sum[K];
int read(){
int su=0,pp=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')pp=-1;ch=getchar();}
while(ch>='0'&&ch<='9'){su=su*10+ch-'0';ch=getchar();}
return su*pp;
}
void build(int &u,int l,int r){
u=(++node_cnt);//动态开点
if(l==r)return;//叶子
int mid=(l+r)>>1;//取中间值切半
build(ls[u],l,mid);//左儿子
build(rs[u],mid+1,r);//右儿子
return;//处理完毕
}
int update(int u,int l,int r,int k){
int v=(++node_cnt);//开新点
ls[v]=ls[u],rs[v]=rs[u];sum[v]=sum[u]+1;//更新值
if(l==r)return v;//叶子
int mid=(l+r)>>1;//取中间值切半
if(k<=mid)ls[v]=update(ls[v],l,mid,k);//左儿子
else rs[v]=update(rs[v],mid+1,r,k);//右儿子
return v;//处理完毕
}
int query(int u,int v,int l,int r,int k){
int mid=((l+r)>>1);//取中间值切半
int x=sum[ls[v]]-sum[ls[u]];//方向判断
if(l==r)return l;//叶子
if(x>=k)return query(ls[u],ls[v],l,mid,k);//往左跑
else return query(rs[u],rs[v],mid+1,r,k-x);//往右跑
}
int main(){
n=read(),Q=read();//读入
for(int i=1;i<=n;i++)p[i]=read(),a[i]=p[i];//读入
sort(a+1,a+n+1);//排序
m=unique(a+1,a+n+1)-a-1;//去重
build(Rt[0],1,m);//建树
for(int i=1;i<=n;i++){
int tmp=lower_bound(a+1,a+m+1,p[i])-a;
Rt[i]=update(Rt[i-1],1,m,tmp);
}//预处理点修改
while(Q--){//处理查询
int l=read(),r=read(),k=read();//读入
int id=query(Rt[l-1],Rt[r],1,m,k);//查询
cout<<a[id]<<"\n";//输出
}
return 0;
}
更多运用
静态查询区间内数字小于 \(k\) 的个数
由于本质和上面提到的【静态查询第 \(k\) 小】是类似的,所以不过多叙述,直接放代码啦。
#include<bits/stdc++.h>
#define LL long long
#define UInt unsigned int
#define ULL unsigned long long
#define LD long double
#define pii pair<int,int>
#define pLL pair<LL,LL>
#define pDD pair<LD,LD>
#define fr first
#define se second
#define pb push_back
#define isr insert
using namespace std;
const int N = 1e5+5;
const int K = (N<<5);
struct qry{int l,r,k;}q[N];
int T,n,m,Q,p[N],a[N];
int node_cnt,num_cnt;
int Rt[N],ls[K],rs[K],sum[K];
map<int,int> Ls;
int read(){
int su=0,pp=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')pp=-1;ch=getchar();}
while(ch>='0'&&ch<='9'){su=su*10+ch-'0';ch=getchar();}
return su*pp;
}
void Clear(){
for(int i=1;i<=node_cnt;i++)
ls[i]=0,rs[i]=0,sum[i]=0;
for(int i=1;i<=n;i++)Rt[i]=0;
node_cnt=0,num_cnt=0;
m=0;Ls.clear();return;
}
void build(int &u,int l,int r){
u=(++node_cnt);//动态开点
if(l==r)return;//叶子
int mid=(l+r)>>1;//取中间值切半
build(ls[u],l,mid);//左儿子
build(rs[u],mid+1,r);//右儿子
return;//处理完毕
}
void update(int &u,int v,int l,int r,int k){
u=(++node_cnt);//开新点
ls[u]=ls[v],rs[u]=rs[v],sum[u]=sum[v]+1;//更新值
if(l==r)return;//叶子
int mid=(l+r)>>1;//取中间值切半
if(k<=mid)update(ls[u],ls[v],l,mid,k);//左儿子
else update(rs[u],rs[v],mid+1,r,k);//右儿子
}
int query(int u,int v,int l,int r,int k){
int mid=((l+r)>>1);//取中间值切半
if(l==r)return sum[v]-sum[u];//叶子
if(mid>=k)return query(ls[u],ls[v],l,mid,k);//往左跑
else return query(rs[u],rs[v],mid+1,r,k)+sum[ls[v]]-sum[ls[u]];//往右跑
}
int main(){
T=read();
for(int Case=1;Case<=T;Case++){
n=read(),Q=read();Clear();
for(int i=1;i<=n;i++)
a[i]=read(),Ls[a[i]]=0;
for(int i=1;i<=Q;i++){
q[i].l=read()+1,q[i].r=read()+1;
q[i].k=read();Ls[q[i].k]=0;
}for(auto &u:Ls)u.se=(++num_cnt);
for(int i=1;i<=n;i++)a[i]=Ls[a[i]];
for(int i=1;i<=Q;i++)q[i].k=Ls[q[i].k];
cout<<"Case "<<Case<<":\n";
for(int i=1;i<=n;i++)
update(Rt[i],Rt[i-1],1,num_cnt,a[i]);
for(int i=1;i<=Q;i++){
auto [l,r,k]=q[i];
int res=query(Rt[l-1],Rt[r],1,num_cnt,k);
cout<<res<<"\n";
}
}
return 0;
}
单点修改查询第 \(k\) 小
这个就是常见的“树状数组套主席树”了,由于需要修改,又是单点修改,这恰好是树状数组最擅长的领域!我们只需要把树状数组套上来,把需要更改的“链”的根节点存在一个临时数组里,修改的时候深入到每个需要变动的叶子节点,把信息调整正确后往上传递就可以啦 > < 整体还是非常简单哒!
#include<bits/stdc++.h>
#define LL long long
#define UInt unsigned int
#define ULL unsigned long long
#define LD long double
#define pii pair<int,int>
#define pLL pair<LL,LL>
#define pDD pair<LD,LD>
#define fr first
#define se second
#define pb push_back
#define isr insert
using namespace std;
const int N = 2e5+5;
const int K = N*400;
struct qry{int opt,l,r,k;}q[N];
struct SegT{int sum,ls,rs;}t[K];
int n,m,Q,len,a[N],Ls[N];
int node_cnt,Rt[N],tmr[2][25],cnt[2];
int read(){
int su=0,pp=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')pp=-1;ch=getchar();}
while(ch>='0'&&ch<='9'){su=su*10+ch-'0';ch=getchar();}
return su*pp;
}
void update(int &u,int l,int r,int p,int k){
if(!u)u=(++node_cnt);
t[u].sum+=k;if(l==r)return;
int mid=(l+r)>>1;
if(p<=mid)update(t[u].ls,l,mid,p,k);
else update(t[u].rs,mid+1,r,p,k);
}
void Pre_upd(int x,int val){
int p=lower_bound(Ls+1,Ls+len+1,a[x])-Ls;
for(;x<=n;x+=x&-x)update(Rt[x],1,len,p,val);return;
}
int Ask_query(int l,int r,int k){
if(l==r)return l;
int mid=(l+r)>>1,res=0;
for(int i=1;i<=cnt[1];i++)res+=t[t[tmr[1][i]].ls].sum;
for(int i=1;i<=cnt[0];i++)res-=t[t[tmr[0][i]].ls].sum;
if(k<=res){
for(int i=1;i<=cnt[1];i++)tmr[1][i]=t[tmr[1][i]].ls;
for(int i=1;i<=cnt[0];i++)tmr[0][i]=t[tmr[0][i]].ls;
return Ask_query(l,mid,k);
}else{
for(int i=1;i<=cnt[1];i++)tmr[1][i]=t[tmr[1][i]].rs;
for(int i=1;i<=cnt[0];i++)tmr[0][i]=t[tmr[0][i]].rs;
return Ask_query(mid+1,r,k-res);
}
}
int Pre_query(int l,int r,int k){
memset(tmr,0,sizeof(tmr));
cnt[0]=0,cnt[1]=0;
for(int i=r;i;i-=i&-i)tmr[1][++cnt[1]]=Rt[i];
for(int i=l-1;i;i-=i&-i)tmr[0][++cnt[0]]=Rt[i];
return Ask_query(1,len,k);
}
int main(){
n=read(),Q=read();
for(int i=1;i<=n;i++)
a[i]=read(),Ls[++len]=a[i];
for(int i=1;i<=Q;i++){
char opt;cin>>opt;
if(opt=='Q')q[i].opt=1,q[i].l=read(),q[i].r=read(),q[i].k=read();
else q[i].opt=0,q[i].l=read(),q[i].k=read(),Ls[++len]=q[i].k;
}sort(Ls+1,Ls+len+1);
len=unique(Ls+1,Ls+len+1)-Ls-1;
for(int i=1;i<=n;i++)Pre_upd(i,1);
for(int i=1;i<=Q;i++){
auto [opt,l,r,k]=q[i];
if(opt)cout<<Ls[Pre_query(l,r,k)]<<"\n";
else Pre_upd(l,-1),a[l]=k,Pre_upd(l,1);
}
return 0;
}
求中位数
求中位数嘛……其实没难度,你只需要知道区间的长度然后利用【查询区间第 \(k\) 小】的方式做就可以了。但是有一个好玩的题,它结合了二分、前缀后缀 \(\max\) 的求解,在卡死一个区间的情况下求出最优的中位数大小,是一个非常有意思的题目哦。
放个代码。
#include<bits/stdc++.h>
#define LL long long
#define UInt unsigned int
#define ULL unsigned long long
#define LD long double
#define pii pair<int,int>
#define pLL pair<LL,LL>
#define pDD pair<LD,LD>
#define fr first
#define se second
#define pb push_back
#define isr insert
using namespace std;
const int N = 3e4+5;
const int K = (N<<6);
struct Tree{int ls,rs,sum,lmx,rmx;}t[K];
struct qry{int l,r,k;}q[N];
int T,n,m,Q,p[N],a[N],Ans;
int num_cnt,bel[N],node_cnt,Rt[N];
map<int,int> Ls;
vector<int> num[N];
int read(){
int su=0,pp=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')pp=-1;ch=getchar();}
while(ch>='0'&&ch<='9'){su=su*10+ch-'0';ch=getchar();}
return su*pp;
}
void push_up(int u){
t[u].lmx=max(t[t[u].ls].lmx,t[t[u].ls].sum+t[t[u].rs].lmx);
t[u].rmx=max(t[t[u].rs].rmx,t[t[u].rs].sum+t[t[u].ls].rmx);
t[u].sum=t[t[u].ls].sum+t[t[u].rs].sum;return;
}
void update(int &u,int l,int r,int p,int k){
t[++node_cnt]=t[u],u=node_cnt;
if(l>p||r<p)return;
if(l==r&&p==l){
t[u].sum=k;
if(k>0)t[u].lmx=k,t[u].rmx=k;
return;
}int mid=(l+r)>>1;
update(t[u].ls,l,mid,p,k);
update(t[u].rs,mid+1,r,p,k);
push_up(u);return;
}
int query_sum(int u,int l,int r,int L,int R){
if(r<L||R<l)return 0;
if(L<=l&&r<=R)return t[u].sum;
int mid=(l+r)>>1,res=0;
res+=query_sum(t[u].ls,l,mid,L,R);
res+=query_sum(t[u].rs,mid+1,r,L,R);
return res;
}
int query_lmx(int u,int l,int r,int L,int R){
if(r<L||R<l)return 0;
if(L<=l&&r<=R)return t[u].lmx;
int mid=(l+r)>>1,res=0;
res=max(res,query_lmx(t[u].ls,l,mid,L,R));
res=max(res,query_sum(t[u].ls,l,mid,L,R)+query_lmx(t[u].rs,mid+1,r,L,R));
return res;
}
int query_rmx(int u,int l,int r,int L,int R){
if(r<L||R<l)return 0;
if(L<=l&&r<=R)return t[u].rmx;
int mid=(l+r)>>1,res=0;
res=max(res,query_rmx(t[u].rs,mid+1,r,L,R));
res=max(res,query_sum(t[u].rs,mid+1,r,L,R)+query_rmx(t[u].ls,l,mid,L,R));
return res;
}
int main(){
n=read();
for(int i=1;i<=n;i++)
a[i]=read(),Ls[a[i]]=0;
Q=read();
for(auto &u:Ls)u.se=(++num_cnt),bel[u.se]=u.fr;
for(int i=1;i<=n;i++)
a[i]=Ls[a[i]],num[a[i]].pb(i);
for(int i=1;i<=n;i++)
update(Rt[num_cnt+1],1,n,i,-1);
for(int i=num_cnt;i>=1;i--){
Rt[i]=Rt[i+1];
for(int x:num[i])update(Rt[i],1,n,x,1);
}while(Q--){
int a=read(),b=read(),c=read(),d=read();
int tmp[4]={(a+Ans)%n,(b+Ans)%n,(c+Ans)%n,(d+Ans)%n};
sort(tmp,tmp+4);
a=tmp[0]+1,b=tmp[1]+1,c=tmp[2]+1,d=tmp[3]+1;
int l=1,r=num_cnt,res=1;
while(l<=r){
int mid=(l+r)>>1;
int tsum=query_sum(Rt[mid],1,n,b,c);
int trmx=query_rmx(Rt[mid],1,n,a,b-1);
int tlmx=query_lmx(Rt[mid],1,n,c+1,d);
int tans=tsum+tlmx+trmx;
if(tans>=0)res=mid,l=mid+1;else r=mid-1;
}Ans=bel[res];cout<<Ans<<"\n";
}
return 0;
}
习题:【模板】树套树
其实不算是主席树了,因为它不再是权值线段树,但依然是动态开点可持久化线段树。套了树状数组,题意很简单,但实现非常复杂,因为需要顾及区间第 \(k\) 小、区间数字 \(k\) 的排名、每个数字的前驱后继……等等等等。逻辑还是简单的,那就最后放个代码叭!
#include<bits/stdc++.h>
#define LL long long
#define UInt unsigned int
#define ULL unsigned long long
#define LD long double
#define pii pair<int,int>
#define pLL pair<LL,LL>
#define pDD pair<LD,LD>
#define fr first
#define se second
#define pb push_back
#define isr insert
using namespace std;
const int N = 5e4+5;
const int M = N*150;
struct Query{int opt,l,r,x,k;}q[N];
struct Tree{int sum,ls,rs;}t[M];
int n,Q,a[N],num_cnt,bel[2*N];
int node_cnt,Rt[N],tmp[2][N];
map<int,int> Ls;
int read(){
int su=0,pp=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')pp=-1;ch=getchar();}
while(ch>='0'&&ch<='9'){su=su*10+ch-'0';ch=getchar();}
return su*pp;
}
void push_up(int u){
t[u].sum=t[t[u].ls].sum+t[t[u].rs].sum;return;
}
void change(int &u,int l,int r,int p,int k){
if(!u)u=(++node_cnt);
if(l==r&&l==p){t[u].sum+=k;return;}
int mid=(l+r)>>1;
if(p<=mid)change(t[u].ls,l,mid,p,k);
else change(t[u].rs,mid+1,r,p,k);
push_up(u);return;
}
void add(int u,int k){
int x=a[u];
while(u<=n)
change(Rt[u],1,num_cnt,x,k),u+=u&-u;return;
}
int queryNum(int l,int r,int k){
if(l==r)return l;
int mid=(l+r)>>1,now=0;
for(int i=1;i<=tmp[0][0];i++)now+=t[t[tmp[0][i]].ls].sum;
for(int i=1;i<=tmp[1][0];i++)now-=t[t[tmp[1][i]].ls].sum;
if(k<=now){
for(int i=1;i<=tmp[0][0];i++)tmp[0][i]=t[tmp[0][i]].ls;
for(int i=1;i<=tmp[1][0];i++)tmp[1][i]=t[tmp[1][i]].ls;
return queryNum(l,mid,k);
}else{
for(int i=1;i<=tmp[0][0];i++)tmp[0][i]=t[tmp[0][i]].rs;
for(int i=1;i<=tmp[1][0];i++)tmp[1][i]=t[tmp[1][i]].rs;
return queryNum(mid+1,r,k-now);
}
}
int askNum(int l,int r,int k){
tmp[0][0]=0,tmp[1][0]=0;l--;
while(r)tmp[0][++tmp[0][0]]=Rt[r],r-=r&-r;
while(l)tmp[1][++tmp[1][0]]=Rt[l],l-=l&-l;
return queryNum(1,num_cnt,k);
}
int queryRank(int l,int r,int k){
if(l==r)return 0;
int mid=(l+r)>>1;
if(k<=mid){
for(int i=1;i<=tmp[0][0];i++)tmp[0][i]=t[tmp[0][i]].ls;
for(int i=1;i<=tmp[1][0];i++)tmp[1][i]=t[tmp[1][i]].ls;
return queryRank(l,mid,k);
}else{
int now=0;
for(int i=1;i<=tmp[0][0];i++)
now+=t[t[tmp[0][i]].ls].sum,
tmp[0][i]=t[tmp[0][i]].rs;
for(int i=1;i<=tmp[1][0];i++)
now-=t[t[tmp[1][i]].ls].sum,
tmp[1][i]=t[tmp[1][i]].rs;
return now+queryRank(mid+1,r,k);
}
}
int askRank(int l,int r,int k){
tmp[0][0]=0,tmp[1][0]=0;l--;
while(r)tmp[0][++tmp[0][0]]=Rt[r],r-=r&-r;
while(l)tmp[1][++tmp[1][0]]=Rt[l],l-=l&-l;
return queryRank(1,num_cnt,k)+1;
}
int Find_pre(int l,int r,int k){
int rk=askRank(l,r,k)-1;
if(!rk)return 0;
else return askNum(l,r,rk);
}
int Find_nxt(int l,int r,int k){
if(k==num_cnt)return num_cnt+1;
int rk=askRank(l,r,k+1);
if(rk==r-l+2)return num_cnt+1;
else return askNum(l,r,rk);
}
int main(){
n=read(),Q=read();
for(int i=1;i<=n;i++)a[i]=read(),Ls[a[i]]=0;
for(int i=1;i<=Q;i++){
q[i].opt=read();
if(q[i].opt==3)q[i].x=read(),q[i].k=read();
else q[i].l=read(),q[i].r=read(),q[i].k=read();
if(q[i].opt!=2)Ls[q[i].k]=0;
}for(auto &u:Ls)u.se=(++num_cnt),bel[u.se]=u.fr;
bel[0]=-2147483647,bel[num_cnt+1]=2147483647;
for(int i=1;i<=n;i++)a[i]=Ls[a[i]],add(i,1);
for(int i=1;i<=Q;i++)
if(q[i].opt!=2)q[i].k=Ls[q[i].k];
for(int i=1;i<=Q;i++)
if(q[i].opt==1)cout<<askRank(q[i].l,q[i].r,q[i].k)<<"\n";
else if(q[i].opt==2)cout<<bel[askNum(q[i].l,q[i].r,q[i].k)]<<"\n";
else if(q[i].opt==3)add(q[i].x,-1),a[q[i].x]=q[i].k,add(q[i].x,1);
else if(q[i].opt==4)cout<<bel[Find_pre(q[i].l,q[i].r,q[i].k)]<<"\n";
else cout<<bel[Find_nxt(q[i].l,q[i].r,q[i].k)]<<"\n";
return 0;
}
另外的,这题好像可以用分块大法水过去哦!
概括与总结
主席树,即可持久化权值线段树,是处理静态区间第 \(k\) 小、前驱后继等查询的便利工具。其核心思想在于保存历史版本,通过动态开点避免空间爆炸,并利用前缀和思想实现任意区间查询。从静态查询到单点修改(树状数组套主席树),再到中位数、排名等复杂操作,主席树展现了强大的扩展性。尽管其代码较长,但可以干脆利落地解决不少难题,是一个非常棒的工具哦!
Thanks reading.

浙公网安备 33010602011771号