和题解1的思路一样但是wa了。。。
#include<iostream>
#include<cstring>
#include<vector>
#include<algorithm>
using namespace std;
typedef long long ll;
const int N=1e6+10;
int n,cot;
int head[N],f[N][20],dep[N],siz[N];
int son[N];
vector<int> col[N];
struct edge
{
int to,nxt;
}eg[N<<1];
void add(int a,int b)
{
eg[++cot].to=b;
eg[cot].nxt=head[a];
head[a]=cot;
}
void dfs(int u,int fa)
{
siz[u]=1;
for(int i=head[u];i;i=eg[i].nxt)
{
int v=eg[i].to;
if(v==fa) continue;
f[v][0]=u;
dep[v]=dep[u]+1;
dfs(v,u);
siz[u]+=siz[v];
}
}
void _init()
{
for(int j=1;j<=18;j++)
{
for(int i=1;i<=n;i++)
{
f[i][j]=f[f[i][j-1]][j-1];
}
}
}
int lca(int u,int v)
{
if(dep[u]<dep[v]) swap(u,v);
for(int i=18;i>=0;i--)
{
if(dep[f[u][i]]>=dep[v]) u=f[u][i];
}
if(u==v) return u;
for(int i=18;i>=0;i--)
{
if(f[u][i]!=f[v][i])
{
u=f[u][i];
v=f[v][i];
}
}
return f[u][0];
}
bool cmp(int a,int b)
{
return dep[a]<dep[b];
}
int main()
{
scanf("%d",&n);
for(int i=1;i<=n;i++)
{
int w;
scanf("%d",&w);
col[w].push_back(i);
}
for(int i=1;i<n;i++)
{
int u,v;
scanf("%d%d",&u,&v);
add(u,v);
add(v,u);
}
dfs(1,0);
_init();
for(int i=1;i<=n;i++)
{
ll ans=0;
int sz=col[i].size();
if(sz==0) ans=(ll)n*(n-1)/2;
else if(sz==1)
{
int u=col[i][0];
ans=siz[u]*(n-siz[u]+1)-1;
int idx=0;
for(int j=head[u];j;j=eg[j].nxt)
{
if(eg[j].to!=f[u][0]) son[++idx]=eg[j].to;
}
for(int j=1;j<=idx;j++)
{
for(int k=j+1;k<=idx;k++)
{
ans+=(ll)siz[son[j]]*siz[son[k]];
}
}
}
else
{
bool flag=0;
int pos;
for(int j=0;j<sz;j++)
{
son[j+1]=col[i][j];
}
sort(son+1,son+sz+1,cmp);
for(int j=sz-1;j>0;j--)
{
int lc=lca(son[j],son[sz]);
if(son[j]!=lc)
{
pos=son[j];
flag=1;
break;
}
}
if(!flag)
{
int u=son[2];
if(f[u][0]!=son[1])
{
for(int j=18;j>=0;j--)
{
if(dep[f[u][i]]>dep[son[1]]) u=f[u][i];
}
}
ans=(ll)siz[son[sz]]*(n-siz[u]);
}
else
{
int lc=lca(pos,son[sz]);
for(int j=1;j<=sz;j++)
{
int lc1=lca(son[j],son[sz]),lc2=lca(son[j],pos);
if(son[j]==lc1||son[j]==lc2)
{
if(dep[son[j]]<dep[lc])
{
flag=0;
break;
}
}
else
{
flag=0;
break;
}
}
ans=flag? (ll)siz[pos]*siz[son[sz]]:0;
}
}
printf("%lld\n",ans);
}
return 0;
}