他一直WA on #10
#include <bits/stdc++.h>
using namespace std;
#define N 114514 * 2
#define M 1919810
#define fi first
#define se second
typedef long long ll;
typedef pair<int, int> pii;
template<typename T> inline T read() {
T x = 0, f = 1; char ch = 0;
for(; !isdigit(ch); ch = getchar()) if(ch == '-') f = -1;
for(; isdigit(ch); ch = getchar()) x = (x << 3) + (x << 1) + (ch - '0');
return x * f;
}
template<typename T> inline void print(T x) {
if(x < 0) putchar('-'), x = -x;
if(x > 9) print(x / 10);
putchar(x % 10 + '0');
}
struct edge {
int to, nxt;
} e[N * 2];
int head[N], tot;
void add(int u, int v) {
e[++ tot] = {v, head[u]};
head[u] = tot;
}
struct query {
int l, r, id, p, lca;
bool operator < (const query &other) const {
if(p != other.p) return p < other.p;
return p & 1 ? r < other.r : r > other.r;
}
} a[N];
int n, m, blk, c[N], col, cnt[N], sum, cur[N], ans[N];
unordered_map<int, int> mp;
int fa[N][30], idx, st[N], ed[N], dep[N], bel[N];
void dfs(int u, int f) {
fa[u][0] = f;
dep[u] = dep[f] + 1;
st[u] = ++ idx, bel[idx] = u;
for(int i = head[u]; i; i = e[i].nxt) {
int v = e[i].to;
if(v == f) continue;
dfs(v, u);
}
ed[u] = ++ idx, bel[idx] = u;
}
int lca(int u, int v) {
if(dep[u] > dep[v]) swap(u, v);
while(dep[u] < dep[v]) {
for(int i = 22; i >= 0; -- i)
if(dep[fa[v][i]] >= dep[u]) v = fa[v][i];
}
if(u == v) return u;
while(fa[u][0] != fa[v][0]) {
for(int i = 22; i >= 0; -- i)
if(fa[u][i] != fa[v][i]) u = fa[u][i], v = fa[v][i];
}
return fa[u][0];
}
void add(int x) {
sum += (++ cnt[x] == 1);
}
void del(int x) {
sum -= (-- cnt[x] == 0);
}
void opt(int x) {
cur[x] ? del(c[x]) : add(c[x]);
cur[x] ^= 1;
}
int main() {
n = read<int>(), m = read<int>();
blk = sqrt(n);
for(int i = 1; i <= n; ++ i) {
int x = read<int>();
if(!mp[x]) mp[x] = ++ col;
c[i] = mp[x];
}
for(int i = 1; i < n; ++ i) {
int u = read<int>(), v = read<int>();
add(u, v), add(v, u);
}
dfs(1, 1);
for(int j = 1; j <= 22; ++ j)
for(int i = 1; i <= n; ++ i)
fa[i][j] = fa[fa[i][j - 1]][j - 1];
for(int i = 1; i <= m; ++ i) {
int u = read<int>(), v = read<int>();
if(st[u] > st[v]) swap(u, v);
int anc = lca(u, v);
if(anc == u) {
a[i].l = st[u], a[i].r = st[v];
a[i].p = a[i].l / blk, a[i].id = i;
a[i].lca = 0;
} else {
a[i].l = ed[u], a[i].r = st[v];
a[i].p = a[i].l / blk, a[i].id = i;
a[i].lca = anc;
}
}
sort(a + 1, a + 1 + m);
for(int i = 1, l = 1, r = 0; i <= m; ++ i) {
while(l > a[i].l) opt(bel[-- l]);
while(r < a[i].r) opt(bel[++ r]);
while(l < a[i].l) opt(bel[l ++]);
while(r > a[i].r) opt(bel[r --]);
if(a[i].lca) opt(bel[a[i].lca]);
ans[a[i].id] = sum;
if(a[i].lca) opt(bel[a[i].lca]);
}
for(int i = 1; i <= m; ++ i) print(ans[i]), putchar('\n');
return 0;
}