蒟蒻求助并查集板子题,WA on #10
查看原帖
蒟蒻求助并查集板子题,WA on #10
310773
PCCP楼主2023/7/27 15:38

RT,感觉我的思路和题解差不多,但是在第十个点错了,不知道是哪里的问题,调了快一天了,特来求助各位大佬。

悬赏一关注,代码如下:

#include<iostream>
#include<cstring>
#include<cmath>
#include<cstdio>
#include<algorithm>
#include<queue>
#include<set>
#include<vector>
#include<map>
using namespace std;
typedef pair<int,int> PII;
const int N=4e5+10;
const int M=8e5+10;
int n,k,r,m,fa[N],siz[N],rt;
int he[N],ne[M<<1],to[M<<1],tot=1;
int fat[30][N],dep[N];
vector<int> sx,sy;
bool st[N];
inline void addedge(int x,int y){
	to[++tot]=y;
	ne[tot]=he[x];
	he[x]=tot;
}
inline int find(int x){
	if(fa[x]==x){
		return x;
	}
	return fa[x]=find(fa[x]);
}
inline void unify(int x,int y){
	int fx=find(x),fy=find(y);
	if(fx==fy){
		return;
	}
	if(siz[fx]<siz[fy]){
		swap(fx,fy);
	}
	fa[fy]=fx;
	siz[fx]+=siz[fy];
}
inline void spread(int root,int x,int f,int step){
	if(step>k/2){
		return;
	}
	if(fa[x]!=x||st[x]==true){
		unify(root,x);
		return;
	}
	unify(root,x);
	for(int i=he[x];i;i=ne[i]){
		int v=to[i];
		if(v!=f){
			spread(root,v,x,step+1);
		}
	}
}
inline void dfs(int x,int f,int deep){
	dep[x]=deep;
	if(st[x]==true){
		for(int i=he[x];i;i=ne[i]){
			int v=to[i];
			spread(x,v,x,1);
		}
	}
	for(int i=he[x];i;i=ne[i]){
		int v=to[i];
		if(v==f){
			continue;
		}
		dfs(v,x,deep+1);
		fat[0][v]=x;
	}
}
inline int lca(int x,int y){
	int l=0;
	while((1<<l)<=n){
		l++;
	}
	if(dep[x]<dep[y]){
		swap(x,y);
	}
	for(int i=l;i>=0;i--){
		if(dep[y]<=dep[x]-(1<<i)){
			x=fat[i][x];
		}
	}
	if(x==y){
		return x;
	}
	for(int i=l;i>=0;i--){
		if(fat[i][x]!=fat[i][y]){
			x=fat[i][x];
			y=fat[i][y];
		}
	}
	return fat[0][x];
}
inline getori(int x,int len){
	for(int i=30;i>=0;i--){
		if((1<<i)<=len){
			x=fat[i][x];
			len-=(1<<i);
		}
	}
	return x;
}
int main(){
	scanf("%d%d%d",&n,&k,&r);
	k*=2;
	int x,y,z;
	for(int i=1;i<n;i++){
		scanf("%d%d",&x,&y);
		addedge(x,n+i);
		addedge(n+i,x);
		addedge(n+i,y);
		addedge(y,n+i);
	}
	for(int i=1;i<=r;i++){
		scanf("%d",&x);
		st[x]=true;
		rt=x;
	}
	for(int i=1;i<=2*n;i++){
		fa[i]=i;
		siz[i]=1;
	}
	dfs(1,0,0);
	for(int i=1;(1<<i)<=2*n;i++){
		for(int j=1;j<=2*n;j++){
			fat[i][j]=fat[i-1][fat[i-1][j]];
		}
	}
	scanf("%d",&m);
	while(m--){
		scanf("%d%d",&x,&y);
		if(dep[x]>dep[y]){
			swap(x,y);
		}
		if(find(x)==find(y)){
			printf("YES\n");
			continue;
		}
		int LCA=lca(x,y);
		if(dep[x]+dep[y]-2*dep[LCA]<=k){
			printf("YES\n");
			continue;
		}
		int lx,ly,len;
		if(dep[x]-dep[LCA]<k/2){
			len=dep[y]-dep[LCA]-(k/2-(dep[x]-dep[LCA]));
			sx.push_back(getori(y,len));
		}
		if(dep[y]-dep[LCA]<k/2){
			len=dep[x]-dep[LCA]-(k/2-(dep[y]-dep[LCA]));
			sy.push_back(getori(x,len));
		}
		len=k/2;
		sx.push_back(getori(x,len));
		sy.push_back(getori(y,len));
		lx=sx.size(),ly=sy.size();
		for(int i=0;i<lx;i++){
			for(int j=0;j<ly;j++){
				if(find(sx[i])==find(sy[j])){
					goto p1;
				}
			}
		}
		printf("NO\n");
		sx.clear();
		sy.clear();
		continue;
		p1 : ;
		printf("YES\n");
		sx.clear();
		sy.clear();
	}
}
2023/7/27 15:38
加载中...