求助,20pts,除5,6 ac外全部re
查看原帖
求助,20pts,除5,6 ac外全部re
817442
Sakuya_maid楼主2023/8/10 16:24
#include <bits/stdc++.h>

using namespace std;
using LL = long long;

constexpr int N = 6e5 + 5;

int ne[N], e[N], h[N], idx;

int dfsn[N], fa[N], son[N], top[N], timestamp, siz[N], dep[N];

int a[N], rnk[N];

int ans[N];

int n;

void add(int u, int v)
{
    e[idx] = v;
    ne[idx] = h[u];
    h[u] = idx ++;
}

void dfs1(int u, int f)
{
    siz[u] = 1;
    fa[u] = f;
    dep[u] = dep[f] + 1;

    int maxn = -1;

    for(int i = h[u]; ~i; i = ne[i])
    {
        int j = e[i];

        if(j == f)continue;

        dfs1(j, u);

        siz[u] += siz[j];

        if(siz[j] > maxn)
        {
            maxn = son[j];
            son[u] = j;
        }
    }
}

void dfs2(int u, int v)
{
    dfsn[u] = ++ timestamp;

    rnk[timestamp] = u;

    top[u] = v;

    if(!son[u])return;

    dfs2(son[u], v);

    for(int i = h[u]; ~i; i = ne[i])
    {
        int j = e[i];

        if(j == son[j] || j == fa[u])continue;

        dfs2(j, j);
    }
}

int lca(int u, int v)
{
    while(top[u] != top[v])
    {
        if(dep[top[u]] < dep[top[v]])u ^= v ^= u ^= v;
        u = fa[top[u]];
    }

    return dep[u] < dep[v] ? u : v;
}

void answer(int u, int f)
{
    for(int i = h[u]; ~i; i = ne[i])
    {
        int j = e[i];

        if(j == f)continue;

        answer(j, u);

        ans[u] += ans[j];
    }
}

void solve()
{
    cin >> n;

    memset(h, -1, sizeof h);

    for(int i = 1; i <= n; ++ i)cin >> a[i];

    for(int i = 1; i < n; ++ i)
    {
        int u, v;
        cin >> u >> v;
        add(u, v);
        add(v, u);
    }

    dfs1(1, 0);
    dfs2(1, 0);

    for(int i = 1; i < n; ++ i)
    {
        int u = a[i], v = a[i + 1];
        int w = lca(u, v);

        ans[u] ++;
        ans[v] ++;
        ans[w] --;
        ans[fa[w]] --;
    }

    answer(1, 0);

    for(int i = 2; i <= n; ++ i)
    ans[a[i]] --;

    for(int i = 1; i <= n; ++ i)
    cout << ans[i] << '\n';
}

signed main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    // int T;
    // for (cin >> T; T -- ; )
        solve();

}
2023/8/10 16:24
加载中...