SPOJ COT2 - Count on a tree II 题解 树上莫队
题目链接:
- https://www.luogu.com.cn/problem/SP10707 (spoj 的 remote judge好像也不行了,但是可以看中文题面)
- 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;
}
浙公网安备 33010602011771号