点分树学习笔记
把点分治过程中依次选取的重心连成一棵树,就是点分树。
考虑维护这样一颗树有什么优点:
- 层数上只有 \(\log(n)\) 层,可以做一些比较暴力的算法。
- 对于带有修改的点分治问题,可以免去进行每次点分治,支持在线修改。
可以发现,点分树中 \(u,v\) 的 LCA 一定在实际树中 \(u \rightarrow v\) 的路径上(该点一定是第一次选取的点将 \(u,v\) 分开到两个子树中,所以只可能在路径上),所以只要找 \(u\) 的所有祖先节点,并将其作为中间点统计点数,就可以得到所有与 \(u\) 距离为 \(k\) 的点。
可以在每个点上建立一个数据结构(BIT 或 SEG 均可)维护距离该节点的子树(点分树)中距离为 \([0,n]\) 的点的点权和,然后遇到一个询问或修改就向上不断跳父亲修改或查询即可。
这样做还有一个问题,一个节点会计算在 \(u\) 所属子树内的点权和,这些路径显然是不合法的,所以在每个点再开一个数据结构,维护其父亲节点到该节点的子树中距离为 \([0,n]\) 的点权和,统计时减去 \(u\) 所属子树的点权和即可。
需要注意倍增 LCA 常数较大,建议换成树剖 LCA 或者欧拉序求 LCA。
CODE
#include<bits/stdc++.h>
#define fst first
#define sec second
#define mkp(a,b) make_pair(a,b)
#define usetime() (double)clock () / CLOCKS_PER_SEC * 1000.0
using namespace std;
typedef long long LL;
typedef pair<int,int> pii;
const int maxn=1e5+5;
void read(int& x){
char c;
bool f=0;
while((c=getchar())<48) f|=(c==45);
x=c-48;
while((c=getchar())>47) x=(x<<3)+(x<<1)+c-48;
x=(f ? -x : x);
}
int n,q;
vector<int> mp[maxn];
bool vis[maxn];
int rt,s[maxn],dfa[maxn],a[maxn];
int dep[maxn];
int fa[maxn], son[maxn], top[maxn], szTree[maxn];
struct SEG{
int t[maxn<<8],ls[maxn<<8],rs[maxn<<8],tot;
SEG(){tot=0;}
void update(int x,int p,int& u,int l=0,int r=n){
if(!u) u=++tot;
if(r<x||l>x) return;
if(l==r){
t[u]+=p; return;
}
int mid=(l+r)>>1;
update(x,p,ls[u],l,mid),update(x,p,rs[u],mid+1,r);
t[u]=t[ls[u]]+t[rs[u]];
}
int query(int L,int R,int u,int l=0,int r=n){
if((!u)||r<L||l>R) return 0;
if(l>=L&&r<=R) return t[u];
int mid=(l+r)>>1;
return query(L,R,ls[u],l,mid)+query(L,R,rs[u],mid+1,r);
}
}t1,t2;
void dfs1(int u,int f){
dep[u]=dep[f]+1;
fa[u]=f;
szTree[u]=1;
son[u]=0;
int maxsz=0;
for(int v : mp[u]){
if(v==f) continue;
dfs1(v,u);
szTree[u]+=szTree[v];
if(szTree[v]>maxsz){
maxsz=szTree[v];
son[u]=v;
}
}
}
void dfs2(int u,int tp){
top[u]=tp;
if(son[u]) dfs2(son[u],tp);
for(int v : mp[u]){
if(v==fa[u] || v==son[u]) continue;
dfs2(v,v);
}
}
int lca(int x,int y){
while(top[x]!=top[y]){
if(dep[top[x]] < dep[top[y]]) swap(x,y);
x=fa[top[x]];
}
return dep[x]<dep[y] ? x : y;
}
int dis(int x,int y){
return dep[x]+dep[y]-2*dep[lca(x,y)];
}
void get_size(int u,int fa){
s[u]=1;
for(int v : mp[u]){
if(v==fa||vis[v]) continue;
get_size(v,u),s[u]+=s[v];
}
}
void find_rt(int u,int fa,int sz){
int mx=0;
for(int v : mp[u]){
if(v==fa||vis[v]) continue;
find_rt(v,u,sz);
mx=max(mx,s[v]);
}
mx=max(mx,sz-s[u]);
if(mx<=sz/2) rt=u;
}
void dfstr(int u,int sz){
vis[u]=1;
for(int v : mp[u]){
if(vis[v]) continue;
get_size(v,u),find_rt(v,u,s[v]);
dfa[rt]=u,dfstr(rt,s[v]);
}
}
int rs1[maxn],rs2[maxn];
void modify(int x,int v){
int now=x;
while(now){
t1.update(dis(x,now),v,rs1[now]);
if(dfa[now]) t2.update(dis(dfa[now],x),v,rs2[now]);
now=dfa[now];
}
}
int query(int x,int k){
int now=x,lst=0,ans=0;
while(now){
if(dis(now,x)>k){
lst=now,now=dfa[now]; continue;
}
ans+=t1.query(0,k-dis(now,x),rs1[now]);
if(lst) ans-=t2.query(0,k-dis(now,x),rs2[lst]);
lst=now,now=dfa[now];
}
return ans;
}
int main(){
read(n),read(q);
for(int i=1;i<=n;i++) read(a[i]);
for(int i=1;i<n;i++){
int u,v; read(u),read(v);
mp[u].push_back(v),mp[v].push_back(u);
}
dfs1(1,0);
dfs2(1,1);
rt=1,get_size(1,1),find_rt(1,1,n),dfstr(rt,n);
for(int i=1;i<=n;i++) modify(i,a[i]);
int lastans=0;
while(q--){
int op,x,y; read(op),read(x),read(y);
x^=lastans,y^=lastans;
if(op==0) printf("%d\n",lastans=query(x,y));
else modify(x,y-a[x]),a[x]=y;
}
return 0;
}

浙公网安备 33010602011771号