点分树学习笔记

把点分治过程中依次选取的重心连成一棵树,就是点分树。

考虑维护这样一颗树有什么优点:

  1. 层数上只有 \(\log(n)\) 层,可以做一些比较暴力的算法。
  2. 对于带有修改的点分治问题,可以免去进行每次点分治,支持在线修改。

例题

可以发现,点分树中 \(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;
}
posted @ 2026-07-21 20:24  huangems  阅读(8)  评论(0)    收藏  举报