将以下代码
#include <bits/stdc++.h>
#define int long long
using namespace std;
int n,m;
int son[100005],siz[100005],fa[100005],top[100005],id[100005],dep[100005],cnt=0;
int ans[100005],a[100005],rid[100005];
vector<int>v[100005],edge[100005];
void dfs1(int u,int f,int deep)
{
siz[u]=1;
fa[u]=f;
dep[u]=deep;
int maxson=-1;
for(int i=0;i<edge[u].size();i++)
{
int v=edge[u][i];
if(v==f)continue;
dfs1(v,u,deep+1);
siz[u]+=siz[v];
if(siz[v]>maxson)
{
maxson=siz[v];
son[u]=v;
}
}
}
void dfs2(int u,int topf)
{
top[u]=topf;
id[u]=++cnt;
rid[cnt]=u;
if(!son[u])return ;
dfs2(son[u],topf);
for(int i=0;i<edge[u].size();i++)
{
int v=edge[u][i];
if(v==fa[u]||v==son[u])continue;
dfs2(v,v);
}
}
signed main()
{
cin>>n>>m;
for(int i=1;i<=n;i++)
cin>>a[i];
for(int i=1;i<n;i++)
{
int u,v;
cin>>u>>v;
edge[u].push_back(v);
edge[v].push_back(u);
}
dfs1(1,0,1);
dfs2(1,1);
for(int i=1;i<=n;i++)
v[a[i]].push_back(id[i]);
while(m--)
{
int a,b,c;
cin>>a>>b>>c;
int flag=0;
while(top[a]!=top[b])
{
if(dep[top[a]]<dep[top[b]])swap(a,b);
vector<int>::iterator it=lower_bound(v[c].begin(),v[c].end(),id[top[a]]);
if(it!=v[c].end()&&*it<=id[a])
{
flag=1;
break;
}
a=fa[top[a]];
}
if(!flag)
{
if(id[a]>id[b])swap(a,b);
vector<int>::iterator it=lower_bound(v[c].begin(),v[c].end(),id[a]);
if(it!=v[c].end()&&*it<=id[b])flag=1;
}
ans[++ans[0]]=flag;
}
for(int i=1;i<=ans[0];i++)
cout<<ans[i];
return 0;
}
中的
for(int i=1;i<=n;i++)
v[a[i]].push_back(id[i]);
改为
for(int i=1;i<=cnt;i++)
v[a[rid[i]]].push_back(i);
就对了 这是为什么?