我感觉我的代码没啥问题,样例也过了,但是WA了
#include<bits/stdc++.h>
using namespace std;
#define int long long
inline int read()
{
int x=0,f=1;char ch=getchar();
while (ch<'0'||ch>'9'){if (ch=='-') f=-1;ch=getchar();}
while (ch>='0'&&ch<='9'){x=x*10+ch-48;ch=getchar();}
return x*f;
}
const int N = 2e6;
int n, m, sq;
int h[N], e[N], ne[N], idx;
bitset<N> st; vector<int> B; int vis[N], ans[N], cnt[N], res, col[N];
int sizes[N], id[N], f[N], g[N], son[N], top[N], dfn, fa[N], dep[N], rev[N];
struct Query
{
int id, l, r;
inline bool operator<(const Query &o) const
{
if (l / sq != o.l / sq) return l < o.l;
if (l / sq & 1) return r < o.r;
return r > o.r;
}
} Q[N];
inline void add(int a, int b)
{
e[idx] = b;
ne[idx] = h[a];
h[a] = idx ++ ;
}
void dfs1(int u, int father)
{
sizes[u] = 1;
dep[u] = dep[father] + 1;
fa[u] = father;
for (int i = h[u]; i != -1; i = ne[i])
{
int j = e[i];
if (j == father) continue;
dfs1(j, u);
sizes[u] += sizes[j];
if (!son[u] || sizes[son[u]] < sizes[j])
son[u] = j;
}
}
void dfs2(int u, int topf)
{
st[u] = true;
top[u] = topf;
id[f[u] = ++ dfn] = u;
if (son[u]) dfs2(son[u], topf);
for (int i = h[u]; i != -1; i = ne[i])
if (!st[e[i]])
dfs2(e[i], e[i]);
id[g[u] = ++ dfn] = u;
}
inline int LCA(int x, int y)
{
while (top[x] != top[y])
dep[top[x]] >= dep[top[y]] ? x = fa[top[x]] : y = fa[top[y]];
return dep[x] < dep[y] ? x : y;
}
inline void add(int x) { res += ( ++ cnt[col[x]] == 1); }
inline void del(int x) { res -= ( -- cnt[col[x]] == 0); }
inline void update(int pos) { (!vis[pos]) ? add(pos) : del(pos); vis[pos] ^= 1; }
signed main()
{
memset(h, -1, sizeof(h));
n = read(), m = read();
sq = pow(n, 2.0 / 3.0);
for (int i = 1; i <= n; i ++ )
col[i] = read(), B.push_back(col[i]);
sort(B.begin(), B.end());
B.erase(unique(B.begin(), B.end()), B.end());
for (int i = 1; i <= n; i ++ )
col[i] = lower_bound(B.begin(), B.end(), col[i]) - B.begin();
for (int i = 1; i <= n - 1; i ++ )
{
int u = read(), v = read();
add(u, v), add(v, u);
}
dfs1(1, 0), dfs2(1, 1);
// for (int i = 1; i <= dfn; i ++ )
// cout << id[i] << " ";
// cout << endl;
for (int i = 1; i <= m; i ++ )
{
int l = read(), r = read();
Q[i].id = i;
Q[i].l = LCA(l, r) == l ? f[l] : g[l];
Q[i].r = f[r];
}
sort(Q + 1, Q + m + 1);
int l = 1, r = 0;
for (int i = 1; i <= m; i ++ )
{
while (l > Q[i].l) update(id[ -- l]);
while (r < Q[i].r) update(id[ ++ r]);
while (l < Q[i].l) update(id[l ++ ]);
while (r > Q[i].r) update(id[r -- ]);
int x = id[l], y = id[r];
int lca = LCA(x, y);
if (x != lca and y != lca)
{
update(lca);
ans[Q[i].id] = res;
update(lca);
}
else ans[Q[i].id] = res;
}
for (int i = 1; i <= m; i ++ )
cout << ans[i] << endl;
return 0;
}