写了2h了,心态有点炸
#include <bits/stdc++.h>
using namespace std;
#define int long long
#define rep(i,a,b) for(int i=a;i<=b;i++)
const int N=1e5+10;
int n,m;
vector<int> g[N];
map<int,int> mp;
vector<int> ls;
int w[N];
int sz[N],S;
int vis[N];
int zx=0,va=1e9;
int ans[N];
int valer[N];
int sets[N];
void getsz(int x,int fa){
sz[x]=1;
for(auto v:g[x]){
if(v==fa||vis[v])continue;
getsz(v,x);
sz[x]+=sz[v];
}
}
void getcentre(int x,int fa){
int mxn=0;
for(auto v:g[x]){
if(v==fa||vis[v])continue;
getcentre(v,x);
mxn=max(mxn,sz[v]);
}
mxn=max(mxn,S-sz[x]);
if(mxn<va){
va=mxn;
zx=x;
}
}
int dis[N];
int sum=0;
void ads(int x,int fa,int sty[]){
if(!sty[w[x]]){valer[w[x]]+=sz[x];sum+=sz[x];
sty[w[x]]=1;}
for(auto v:g[x]){
if(v==fa||vis[v])continue;
ads(v,x,sty);
}
sty[w[x]]=0;
}
void nads(int x,int fa,int sty[]){
if(!sty[w[x]]){valer[w[x]]-=sz[x];sum-=sz[x];
sty[w[x]]=1;}
for(auto v:g[x]){
if(v==fa||vis[v])continue;
nads(v,x,sty);
}
sty[w[x]]=0;
}
int sznw=0;
int centre;
void con(int x,int fa,int sty[]){
int upto=0,goaled=0;
if(!sty[w[x]]){
upto+=valer[w[centre]]-valer[w[x]];
sty[w[x]]=1;
goaled=1;
}
sznw+=upto;
ans[x]+=sznw+sum;
for(auto v:g[x]){
if(v==fa||vis[v])continue;
con(v,x,sets);
}
if(goaled)sty[w[x]]=0;
sznw-=upto;
}
void del(int x,int fa){
valer[w[x]]=0;
for(auto v:g[x]){
if(v==fa||vis[v])continue;
del(v,x);
}
}
void divide(int x){
zx=0,va=1e9;
getsz(x,0);
S=sz[x];
getcentre(x,0);
x=zx;
centre=x;
getsz(x,0);
vis[x]=1;
del(x,0);
sum=0;
for(auto v:g[x]){
if(vis[v])continue;
getsz(v,x);
ads(v,x,sets);
}
ans[x]+=sum+sz[x]-valer[w[x]];
for(auto v:g[x]){
if(vis[v])continue;
nads(v,x,sets);
int tmp=valer[w[x]];
sum-=valer[w[x]];
sum+=sz[x]-sz[v];
valer[w[x]]=sz[x]-sz[v];
con(v,x,sets);
sum-=sz[x]-sz[v];
valer[w[x]]=tmp;
sum+=tmp;
ads(v,x,sets);
}
for(auto v:g[x]){
if(vis[v])continue;
divide(v);
}
}
signed main(){
ios::sync_with_stdio(0);
cin>>n;
rep(i,1,n)cin>>w[i];
rep(i,1,n-1){
int u,v;cin>>u>>v;
g[u].push_back(v);
g[v].push_back(u);
}
divide(1);
rep(i,1,n)cout<<ans[i]<<'\n';
}