rt,看起来很对但是 Wrong Answer
#include<bits/stdc++.h>
#define MAXN 200002
using namespace std;
int n,k,cnt,root,ans=0x7fffffff;
int c[MAXN],siz[MAXN],maxsiz[MAXN],fa[MAXN],bel[MAXN];
bool merged[MAXN],rooted[MAXN],vis[MAXN];
vector<int>g[MAXN];
vector<int>id[MAXN];
vector<int>del;
inline void addedge(int u,int v)
{
g[u].push_back(v);
return;
}
inline void getfa(int u,int rt)
{
bel[u]=rt;
for(auto v:g[u])
{
if(rooted[v]||v==fa[u])continue;
fa[v]=u;
getfa(v,rt);
}
return;
}
inline void getrt(int u,int fa,int tot)
{
maxsiz[u]=0;
siz[u]=1;
for(auto v:g[u])
{
if(rooted[v]||v==fa)continue;
getrt(v,u,tot);
siz[u]+=siz[v];
maxsiz[u]=max(maxsiz[u],siz[v]);
}
maxsiz[u]=max(maxsiz[u],tot-siz[u]);
if(maxsiz[u]<maxsiz[root]||root==0)root=u;
return;
}
inline void solve(int rt)
{
int cnt=0;
getfa(rt,rt);
queue<int>q;
for(auto u:id[c[rt]])
{
if(u!=rt)q.push(u);
if(bel[u]!=rt)goto END;
}
del.push_back(c[rt]);
merged[c[rt]]=true;
while(!q.empty())
{
int u=q.front();q.pop();
if(!merged[c[fa[u]]])
{
cnt++;
merged[c[fa[u]]]=true;
del.push_back(c[fa[u]]);
for(auto v:id[c[fa[u]]])
{
q.push(v);
if(bel[u]!=rt)goto END;
}
}
}
ans=min(ans,cnt);
END:;
for(auto u:del)merged[u]=false;
del.clear();
return;
}
inline void divide(int rt)
{
solve(rt);
rooted[rt]=true;
for(auto u:g[rt])
{
if(rooted[u])continue;
root=0;
getrt(u,rt,siz[u]);
divide(root);
}
return;
}
int main()
{
// freopen("merge.in","r",stdin);
// freopen("merge.out","w",stdout);
ios::sync_with_stdio(false);
cin.tie(0);cout.tie(0);
cin>>n>>k;
for(int i=1,u,v;i<n;i++)
{
cin>>u>>v;
addedge(u,v);addedge(v,u);
}
for(int i=1;i<=n;i++)
{
cin>>c[i];
id[c[i]].push_back(i);
}
getrt(1,0,n);
divide(root);
printf("%d\n",ans);
return 0;
}