「Ynoi Easy Round 2016」这是我自己的发明

样例调过后居然一遍过了?

首先发现的是换根是假的,本质上不带修改,不同的根对于 \(u\) 来说只有 \(2\) 种情况,一种以 \(1\) 为根,一种以 \(u\) 某棵子树 \(p\) 里的点为根,定义 \(u\) 的 dfs 序为 \(dfn_u\)\(u\) 子树中 dfs 序最大为 \(out_u\),则两种情况分别对应 \([dfn_u, out_u]\)\([1, dfn_p) \cup (out_p, n]\) 两种区间计数,所以查询可以用最多 \(4\) 个四元组 \((l1, r1, l2, r2)\) 查询加在一起计算,然后每个四元组又可以像二位前缀和一样拆成 \(4\) 个前缀查询 \((r1, r2), (r1, l2 - 1), (l1 - 1, r2), (l1 - 1, l2 - 1)\) 做容斥,共拆成 \(16\) 个询问,然后就可以用莫队来解决。

这样好像可以过,但是太麻烦了,我们预处理 \(f_i\) 表示查询 \((1, i, 1, n)\) 的答案,那么对于不同的查询我们可以加入 \(f\) 进行优化,令 \(g(l, r, L, R)\) 表示查询 \((l, r, L, R)\) 的答案。

  • \([l, r], [L, R] \Rightarrow g(l, r, L, R)\)
  • \([l, r], [1, L) \cup (R, n] \Rightarrow f_r - f_{l - 1} - g(l, r, L, R)\)
  • \([1, l) \cup (r, n], [1, L) \cup (R, n] \Rightarrow f_n - (f_r - f_{l - 1}) - (f_R - f_{L - 1}) + g(l, r, L, R)\)

这样我们就只需要查 \(g(l, r, L, R)\),就只需要拆成 \(4\) 个询问了。

时间复杂度 \(O(M \sqrt{N})\),空间复杂度 \(O(N + M)\)

/*
address:https://www.luogu.com.cn/problem/P4689
AC 2026/8/6 17:31
*/
#include<bits/stdc++.h>
using namespace std;
typedef pair<int, int> pii;
#define mkp make_pair
typedef long long LL;
const int N = 1e5 + 5, M = 2e6 + 5, B = 320;
int n, m, q, siz;
int a[N], disc[N], t;
vector<int>G[N];
int dfn[N], out[N], rnk[N], cntn;
LL f[N];
int cnt[N];
inline void dfs(int u, int fa) {
    dfn[u] = ++cntn;
    rnk[cntn] = u;
    f[cntn] = f[cntn - 1] + cnt[a[u]];
    for (auto v : G[u])
        if (v != fa) dfs(v, u);
    for (auto& v : G[u]) v = dfn[v];
    out[u] = cntn;
}
struct query {
    int x, y, bel, id;
    bool operator < (const query& o)const {
        if (bel == o.bel) return bel & 1 ? y < o.y : y > o.y;
        return bel < o.bel;
    }
}qry[M];
LL ans[M >> 2];
inline void init() {
    sort(disc + 1, disc + t + 1);
    t = unique(disc + 1, disc + t + 1) - disc - 1;
    for (int i = 1;i <= n;++i) ++cnt[a[i] = lower_bound(disc + 1, disc + t + 1, a[i]) - disc];
    dfs(1, 0);
}
inline void addquery(int l, int r, int L, int R, int id) {
    qry[++q] = { r, R, (r + siz - 1) / siz, id };
    if (l > 1) qry[++q] = { l - 1, R, (l - 1 + siz - 1) / siz, -id };
    if (L > 1) qry[++q] = { r, L - 1, (r + siz - 1) / siz, -id };
    if (l > 1 && L > 1) qry[++q] = { l - 1, L - 1, (l - 1 + siz - 1) / siz, id };
}
int c1[N], c2[N];
LL cur;
inline void solve() {
    sort(qry + 1, qry + q + 1);
    int x = 0, y = 0;
    for (int i = 1;i <= q;++i) {
        while (x < qry[i].x) ++c1[a[rnk[++x]]], cur += c2[a[rnk[x]]];
        while (y < qry[i].y) ++c2[a[rnk[++y]]], cur += c1[a[rnk[y]]];
        while (y > qry[i].y) --c2[a[rnk[y]]], cur -= c1[a[rnk[y--]]];
        while (x > qry[i].x) --c1[a[rnk[x]]], cur -= c2[a[rnk[x--]]];
        ans[abs(qry[i].id)] += qry[i].id > 0 ? cur : -cur;
    }
}
inline void get(int u, int& l, int& r, int rt) {
    if (u == rt) return l = 1, r = n, void(0);
    if (dfn[rt] > dfn[u] && dfn[rt] <= out[u]) {
        int p = rnk[G[u][upper_bound(G[u].begin(), G[u].end(), dfn[rt]) - G[u].begin() - 1]];
        l = dfn[p], r = out[p];
    }
    else l = dfn[u], r = out[u];
}
inline void read(int& x) {
    x = 0;
    char c = getchar();
    while (c < '0' || c > '9') c = getchar();
    while (c >= '0' && c <= '9') x = x * 10 + c - '0', c = getchar();
}
int main() {
    read(n), read(m);
    siz = sqrt(n);
    for (int i = 1;i <= n;++i) read(a[i]), disc[++t] = a[i];
    for (int i = 1;i < n;++i) {
        int u, v;read(u), read(v);
        G[u].push_back(v), G[v].push_back(u);
    }
    init();
    for (int i = 1, rt = 1;i <= m;++i) {
        ans[i] = -1;
        int op;read(op);
        if (op == 1) read(rt);
        else {
            int u, v;read(u), read(v);
            bool inu = dfn[rt] > dfn[u] && dfn[rt] <= out[u], inv = dfn[rt] > dfn[v] && dfn[rt] <= out[v];
            if (inu) swap(u, v), swap(inu, inv);
            int l, r, L, R;
            get(u, l, r, rt), get(v, L, R, rt);
            if (u == rt) ans[i] = inv ? f[n] - (f[R] - f[L - 1]) : f[R] - f[L - 1];
            else if (v == rt) ans[i] = inu ? f[n] - (f[r] - f[l - 1]) : f[r] - f[l - 1];
            else if (!inu && !inv) ans[i] = 0, addquery(l, r, L, R, i);
            else if (!inu) ans[i] = f[r] - f[l - 1], addquery(l, r, L, R, -i);
            else ans[i] = f[n] - (f[r] - f[l - 1]) - (f[R] - f[L - 1]), addquery(l, r, L, R, i);
        }
    }
    solve();
    for (int i = 1;i <= m;++i)
        if (ans[i] != -1) printf("%lld\n", ans[i]);
    return 0;
}
posted @ 2026-08-06 17:54  keysky  阅读(3)  评论(0)    收藏  举报