求助卡常
查看原帖
求助卡常
551803
BPG_ning楼主2023/7/14 17:21
#include<bits/stdc++.h>
using namespace std; 
const int N=4e4+1;
typedef pair<int,int> pii;
typedef bitset<N> bs;
int n,m,qt=2000,b[N];
int cnt,to[N<<1],nxt[N<<1],head[N];
int num,gj[N],dep[N],maxdep[N],vis[N],down[N];
int fa[N],son[N],sz[N],top[N],up[N][6];
int h,st[N];
bs bit[21][21];
struct node{int x,id;}a[N];
bool cmp_x(node a,node b){return a.x<b.x;}
void add(int x,int y){
	to[++cnt]=y;
	nxt[cnt]=head[x];
	head[x]=cnt;
}
void lsh(){
	sort(a+1,a+1+n,cmp_x);
	int num=0;
	for(int i=1;i<=n;i++){
		if(i==1||a[i].x!=a[i-1].x) num++;
		b[a[i].id]=num;
	}
}
void dfs1(int x,int fath){
	dep[x]=dep[fath]+1;
	maxdep[x]=dep[x];
	sz[x]=1; fa[x]=fath;
	for(int i=head[x];i;i=nxt[i]){
		int y=to[i];
		if(y==fath) continue;
		dfs1(y,x);
		maxdep[x]=max(maxdep[x],dep[y]);
		sz[x]+=sz[y];
		if(sz[son[x]]<sz[y]) son[x]=y;
	}
	if(maxdep[x]-dep[x]>=qt){
		gj[++num]=x;
		vis[x]=num;
		maxdep[x]=dep[x];
	}
	return ;
}
void dfs2(int x,int tp){
	top[x]=tp;
	if(son[x]==0) return ;
	dfs2(son[x],tp);
	for(int i=head[x];i;i=nxt[i]){
		int y=to[i];
		if(y==fa[x]||y==son[x]) continue;
		dfs2(y,y);
	}
}
void dfs3(int x,bs tmp){
	int idx=vis[x];
	tmp[b[x]]=1;
	if(idx!=0){
		int idh=vis[st[h]]; up[x][0]=st[h];
//		for(int i=1;i<=5;i++) up[x][i]=up[up[x][i-1]][i-1];
		bit[idx][idh]=bit[idh][idx]=tmp;
		for(int i=h-1;i>=1;i--){
			int idi=vis[st[i]];
			bit[idx][idi]=bit[idi][idx]=(tmp|bit[idh][idi]);
		}
		tmp=0; st[++h]=x;
	}
	tmp[b[x]]=1;
	for(int i=head[x];i;i=nxt[i]){
		int y=to[i];
		if(y==fa[x]) continue;
		dfs3(y,tmp);
	}
	if(idx!=0) h--;
}
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]];
	}
	if(dep[x]>dep[y]) swap(x,y);
	return x;
}
bs qans(int l,int r){
	bs ans=0;
	while(vis[l]==0&&l!=r){
		ans[b[l]]=1;
		l=fa[l];
	}
//	for(int i=5;i>=0;i--){
//		while(dep[up[l][i]]>=dep[r]){
//			ans=(ans|bit[l][up[l][i]]);
//			l=up[l][i];
//		}
//	}
	while(dep[up[l][0]]>=dep[r]){
		ans=(ans|bit[l][up[l][0]]);
		l=up[l][0];
	}
	while(dep[l]>=dep[r]){
		ans[b[l]]=1;
		l=fa[l];
	}
	return ans;
}
int main(){
	ios::sync_with_stdio(false);
	std::cin.tie(0);
	std::cout.tie(0); 
	freopen("nzq.in","r",stdin);
	freopen("nzq.out","w",stdout);
	cin>>n>>m;
	for(int i=1;i<=n;i++)cin>>a[i].x,a[i].id=i;
	lsh();
	for(int i=1;i<n;i++){
		int x,y;
		cin>>x>>y;
		add(x,y); add(y,x);
	}
	dfs1(1,0);
	dfs2(1,0);
	dfs3(1,0);
	int lastans=0;
	while(m--){
		int x,y;
		cin>>x>>y;
		x=x^lastans;
		int l=lca(x,y);
		lastans=(qans(x,l)|qans(y,l)).count();
		cout<<lastans<<'\n';
	}
	return 0;
}

跟题解不同的是我是去倍增跳块,所以加了个log20

2023/7/14 17:21
加载中...