QOJ5017 相等树链

QOJ5017 相等树链

https://qoj.ac/problem/5017

树上路径问题,考虑点分治求解。

枚举分支中心 \(u\)。设 \(T_1\) 上两端点为 \(x,y\)\(T_2\) 上两端点为 \(z,w\)

  1. \(z,w\in P_1(x,u)\)

    首先判断 \(P_1(x,u)\)\(T_2\) 上是否可能是某条链的子集,如果可能就找到两个端点 \(s_x,t_x\)。合法的充要条件是 \(P_1(x,u)\oplus P_1(y,u)=P_2(s_x,u)\oplus P_2(t_x,u)\),即 \(P_1(x,u)\oplus P_2(s_x,u)\oplus P_2(t_x,u)=P_1(y,u)\)

  2. \(z,w\in P_2(x,u)\),类似做即可。

  3. \(z\in P_1(x,u),w\in P_2(w,u),z\ne u,w\ne u\)

    \(z\) 可以等于 \(s_x\)\(t_x\)\(w\) 可以等于 \(s_y\)\(t_y\)。枚举 \(z,w\) 的取值,合法充要条件是 \(P_1(x,u)\oplus P_2(z,u)=P_1(y,u)\oplus P_2(w,u)\)\(z,w\)\(T_2\) 上位于不同子树内。

\(s_x,t_x\) 可以动态加点维护:新加入一个 \(p\),求出三者中心点 \(m\),若 \(m\) 不是三者之一则不合法,否则除去 \(m\) 后剩下的两个点为新端点。而三者中心点为两两 LCA 的异或和。

集合相等可以赋随机值做 xor-hashing。利用 \(O(1)\) LCA、哈希表、\(O(1)\) 树上 \(k\) 级祖先,精细实现可以 \(O(n\log n)\)

int n;
mt19937 seed(350234);
uniform_int_distribution<ull> rnd(0,(ull)(-1));
//uniform_int_distribution<ull> rnd(0,255);
ull W[N]; ll Ans;
int lg[N],TOT;

struct Edge{
	int to,nxt;
};

namespace T2{
	int head[N],tot;
	Edge edge[N<<1];
	ull d[N];
	int dep[N],fa[N][22];
	int dfn[N],tim;
	
	struct STNode{
		int val,frm;
		
		bool operator < (const STNode& tmp)const{
			return val<tmp.val;
		}
	}st[N][22];
	
	void Add(int u,int v){
		edge[++tot]={v,head[u]};
		head[u]=tot;
	}
	
	void dfs(int x,int pr){
		dep[x]=dep[pr]+1;
		fa[x][0]=pr;
		dfn[x]=++tim;
		st[tim][0]={dep[x],x};
		for(int i=1;i<=lg[dep[x]];i++)
			fa[x][i]=fa[fa[x][i-1]][i-1];
		d[x]=d[pr]^W[x];
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(y==pr) continue;
			dfs(y,x);
		}
	}
	
	int Askst(int l,int r){
		int x=lg[r-l+1];
		return min(st[l][x],st[r-(1<<x)+1][x]).frm;
	}
	
	int Asklca(int u,int v){
		if(u==v) return u;
		if(dfn[u]>dfn[v]) swap(u,v);
		return fa[Askst(dfn[u]+1,dfn[v])][0];
	}
	
	int Middle(int x,int y,int z){
		return Asklca(x,y)^Asklca(x,z)^Asklca(y,z);
	}
	
	ull P(int x,int y){
//		printf("LCA(%d,%d)=%d\n",x,y,Asklca(x,y));
		return d[x]^d[y]^W[Asklca(x,y)];
	}
	
	int Jump(int x,int k){
		for(int i=lg[k];i>=0;i--)
			if(k>>i&1) x=fa[x][i];
		return x;
	}
	
	int Belong(int x,int y){
		if(x==y) return x;
		int LCA=Asklca(x,y);
		if(LCA==x) return Jump(y,dep[y]-dep[x]-1);
		else return fa[x][0];
	}
	
	void Init(){
		dfs(1,0);
		for(int j=1;j<=lg[n];j++)
			for(int i=1;i+(1<<j)-1<=n;i++)
				st[i][j]=min(st[i][j-1],st[i+(1<<(j-1))][j-1]);
	}
}

namespace T1{
	int head[N],tot;
	Edge edge[N<<1];
	int siz[N]; ull f[N],g[N],h[N][2];
	int All,Rt,Mn;
	bool vis[N];
	int s[N],t[N];
	gp_hash_table<ull,int> mf,mg,mh[N];
	
	void Add(int u,int v){
		edge[++tot]={v,head[u]};
		head[u]=tot;
	}
	
	void dfs(int x,int pr){
		siz[x]=1;
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(y==pr) continue;
			dfs(y,x);
			siz[x]+=siz[y];
		}
	}
	
	void FindRt(int x,int pr){
		siz[x]=1; int mx=0;
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(y==pr||vis[y]) continue;
			FindRt(y,x);
			siz[x]+=siz[y];
			Ckmax(mx,siz[y]);
		}
		Ckmax(mx,All-siz[x]);
		if(mx<Mn) Mn=mx,Rt=x;
	}
	
	void dfs1(int x,int pr,int rt){
		f[x]=f[pr]^W[x];
		int p=T2::Middle(s[pr],t[pr],x);
		if(p!=s[pr]&&p!=t[pr]&&p!=x){
			s[x]=t[x]=-1;
			return;
		}
		if(p==s[pr]) s[x]=t[pr],t[x]=x;
		else if(p==t[pr]) s[x]=s[pr],t[x]=x;
		else s[x]=s[pr],t[x]=t[pr];
		g[x]=f[x]^T2::P(rt,s[x])^T2::P(rt,t[x]);
		h[x][0]=f[x]^T2::P(rt,s[x]);
		h[x][1]=f[x]^T2::P(rt,t[x]);
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(vis[y]||y==pr) continue;
			dfs1(y,x,rt);
		}
	}
	
	void dfs2(int x,int pr,int rt){
		if(s[x]==-1) return;
		if(mf.find(g[x])!=mf.end()) Ans+=mf[g[x]];
		if(mg.find(f[x])!=mg.end()) Ans+=mg[f[x]];
		if(s[x]!=rt){
			int bl=T2::Belong(rt,s[x]);
			if(mh[0].find(h[x][0])!=mh[0].end()) Ans+=mh[0][h[x][0]];
			if(mh[bl].find(h[x][0])!=mh[bl].end()) Ans-=mh[bl][h[x][0]];
		}
		if(t[x]!=rt){
			int bl=T2::Belong(rt,t[x]);
			if(mh[0].find(h[x][1])!=mh[0].end()) Ans+=mh[0][h[x][1]];
			if(mh[bl].find(h[x][1])!=mh[bl].end()) Ans-=mh[bl][h[x][1]];
		}
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(y==pr||vis[y]) continue;
			dfs2(y,x,rt);
		}
	}
	
	void dfs3(int x,int pr,int rt){
		if(s[x]==-1) return;
		++mf[f[x]],++mg[g[x]];
		if(s[x]!=rt) ++mh[0][h[x][0]],++mh[T2::Belong(rt,s[x])][h[x][0]];
		if(t[x]!=rt) ++mh[0][h[x][1]],++mh[T2::Belong(rt,t[x])][h[x][1]];
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(y==pr||vis[y]) continue;
			dfs3(y,x,rt);
		}
	}
	
	void dfs4(int x,int pr){
		++TOT;
		siz[x]=1;
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(y==pr||vis[y]) continue;
			dfs4(y,x);
			siz[x]+=siz[y];
		}
	}
	
	void Work(int x){
		s[x]=t[x]=x;
		f[x]=g[x]=h[x][0]=h[x][1]=W[x];
		mf.clear(),mg.clear(),mh[x].clear(),mh[0].clear();
		for(int i=T2::head[x];i;i=T2::edge[i].nxt)
			mh[T2::edge[i].to].clear();
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(!vis[y]) dfs1(y,x,x);
		}
		++mf[f[x]];
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(vis[y]) continue;
			dfs2(y,x,x);
			dfs3(y,x,x);
		}
		dfs4(x,0);
	}
	
	void Solve(int x){
		Work(x);
		vis[x]=1;
		for(int i=head[x];i;i=edge[i].nxt){
			int y=edge[i].to;
			if(vis[y]) continue;
			All=siz[y],Mn=IINF,Rt=0;
			FindRt(y,x); Solve(Rt);
		}
	}
}

signed main(){
	read(n);
	for(int i=2;i<=n;i++){
		int x; read(x);
		T1::Add(x,i),T1::Add(i,x);
	}
	for(int i=2;i<=n;i++){
		int x; read(x);
		T2::Add(x,i),T2::Add(i,x);
	}
	for(int i=1;i<=n;i++) W[i]=rnd(seed);
	for(int i=2;i<=n;i++) lg[i]=lg[i>>1]+1;
	T2::Init();
	T1::Solve(1);
	printf("%lld\n",Ans+n);
	cerr<<TOT<<endl;
	return 0;
}
posted @ 2026-04-16 19:23  XP3301_Pipi  阅读(18)  评论(0)    收藏  举报
Title