#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