SPOJ COT2 - Count on a tree II 题解 树上莫队

题目链接:

  1. https://www.luogu.com.cn/problem/SP10707 (spoj 的 remote judge好像也不行了,但是可以看中文题面)
  2. https://www.spoj.com/problems/COT2/ SPOJ官网题面(英文,可提交)

求 欧拉序,得到一个长度为 \(2n\) 的欧拉序列;

然后在 欧拉序 上进行莫队。

示例程序:

#include <bits/stdc++.h>
using namespace std;
const int maxn = 4e4 + 5, maxm = 1e5 + 5;

int n, m, blo, a[maxn], fa[maxn][16], dep[maxn], id[maxn * 2], idx, dfn[maxn], nfd[maxn], ans[maxm], sum;
vector<int> g[maxn];

void lsh() {
    vector<int> v(a+1, a+n+1);
    sort(v.begin(), v.end());
    v.erase(unique(v.begin(), v.end()), v.end());
    for (int i = 1; i <= n; i++)
        a[i] = lower_bound(v.begin(), v.end(), a[i]) - v.begin() + 1;
}

void dfs(int u, int p) {
    id[ dfn[u] = ++idx ] = u;
    dep[u] = dep[p] + 1;
    fa[u][0] = p;
    for (auto v : g[u])
        if (v != p)
            dfs(v, u);
    id[ nfd[u] = ++idx ] = u;
}

int lca(int x, int y) {
    if (dep[x] < dep[y])
        swap(x, y);
    for (int i = 15; i >= 0; i--) {
        int p = fa[x][i];
        if (dep[p] >= dep[y])
            x = p;
    }
    if (x == y)
        return x;
    for (int i = 15; i >= 0; i--)
        if (fa[x][i] != fa[y][i])
            x = fa[x][i], y = fa[y][i];
    return fa[x][0];
}

struct Query {
    int id, l, r;
} query[maxm];

bool vis[maxn];
int cnt[maxn];

void add(int x) {
    int u = id[x];
    if (!vis[u]) {
        if (++cnt[ a[u] ] == 1) sum++;
    }
    else {
        if (--cnt[ a[u] ] == 0) sum--;
    }
    vis[u] ^= 1;
}

int main() {
    scanf("%d%d", &n, &m);
    blo = sqrt(n * 2);
    for (int i = 1; i <= n; i++)
        scanf("%d", a+i);
    lsh();

    for (int i = 1, u, v; i < n; i++) {
        scanf("%d%d", &u, &v);
        g[u].push_back(v);
        g[v].push_back(u);
    }
    dfs(1, 0);
    for (int i = 1; i <= 15; i++)
        for (int u = 1; u <= n; u++)
            fa[u][i] = fa[ fa[u][i-1] ][i-1];
    for (int i = 1, u, v, l, r; i <= m; i++) {
        scanf("%d%d", &u, &v);
        if (dfn[u] > dfn[v])
            swap(u, v);
        int z = lca(u, v);
        if (z != u && z != v)
            query[i] = {i, nfd[u], dfn[v]};
        else
            query[i] = {i, dfn[u], dfn[v]};
    }
    sort(query+1, query+m+1, [](auto a, auto b) {
        if (a.l / blo != b.l / blo)
            return a.l < b.l;
        return (a.l / blo % 2) ? (a.r < b.r) : (a.r > b.r);
      });
    for (int i = 1, l = 1, r = 0; i <= m; i++) {
        while (l < query[i].l) add(l++);
        while (l > query[i].l) add(--l);
        while (r < query[i].r) add(++r);
        while (r > query[i].r) add(r--);
        int u = id[l], v = id[r], z = lca(u, v);
        if (z == u || z == v) {
            ans[ query[i].id ] = sum;
        }
        else {
            add(dfn[z]);
            ans[ query[i].id ] = sum;
            add(dfn[z]);
        }
    }
    for (int i = 1; i <= m; i++)
        printf("%d\n", ans[i]);
    return 0;
}
posted @ 2026-03-05 11:23  quanjun  阅读(8)  评论(0)    收藏  举报