为什么 WA 50
查看原帖
为什么 WA 50
354310
Tnuzy_plzro楼主2023/3/24 20:29

写了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';
}
2023/3/24 20:29
加载中...