点分治求调
查看原帖
点分治求调
401052
Endline楼主2023/7/21 11:29

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;
}
2023/7/21 11:29
加载中...