基本思路是题解的第一种做法。
代码略去缺省源和头文件
namespace Program
{
#define BL(i) (s * (i - 1) + 1)
#define BR(i) (min(len, s * i))
const int MAXN = 40005;
const int MAXM = 305;
int n, m, s, col[MAXN], tot = 1, hd[MAXN], to[MAXN << 1], nxt[MAXN << 1], len, seq[MAXN << 1], pos[MAXN][2], in[MAXN << 1], cnt[2][MAXM][MAXN], st[MAXM][MAXM], tmp[2][MAXN], fa[MAXN], dep[MAXN], si[MAXN], hs[MAXN], top[MAXN], ans;
inline void Link(int u, int v) { return ++tot, nxt[tot] = hd[u], to[tot] = v, hd[u] = tot, void(); }
inline void Discretization()
{
sort(seq + 1, seq + len + 1), len = unique(seq + 1, seq + len + 1) - seq - 1;
for(register int i = 1; i <= n; i++) col[i] = lower_bound(seq + 1, seq + len + 1, col[i]) - seq;
return len = 0, void();
}
inline int LCA(int u, int v)
{
while(top[u] != top[v])
{
if(dep[top[u]] < dep[top[v]]) swap(u, v);
u = fa[top[u]];
}
if(dep[u] < dep[v]) return u;
return v;
}
inline void DFS1(int u, int f)
{
seq[++len] = u, pos[u][0] = len, fa[u] = f, dep[u] = dep[f] + 1, si[u] = 1, hs[u] = 0;
for(register int i = hd[u]; i; i = nxt[i])
{
int v = to[i];
if(v == f) continue;
DFS1(v, u);
si[u] += si[v];
if(si[v] > si[hs[u]]) hs[u] = v;
}
return seq[++len] = u, pos[u][1] = len, void();
}
inline void DFS2(int u, int t)
{
top[u] = t;
if(hs[u]) DFS2(hs[u], t);
for(register int i = hd[u]; i; i = nxt[i])
{
int v = to[i];
if(v == fa[u] || v == hs[u]) continue;
DFS2(v, v);
}
return;
}
inline void Extend(int u, int *tmp0, int *tmp1, int &res)
{
++tmp0[u];
if(tmp0[u] == 1) res += (!tmp1[col[u]]), ++tmp1[col[u]];
else --tmp1[col[u]], res -= (!tmp1[col[u]]);
return;
}
inline void Init()
{
s = (int)sqrt(len);
for(register int i = 1; i <= len; i++)
{
in[i] = (i - 1) / s + 1, ++tmp[0][seq[i]], ++tmp[1][col[seq[i]]];
if(i == BR(in[i])) for(register int j = 1; j <= n; j++) cnt[0][in[i]][j] = tmp[0][j], cnt[1][in[i]][j] = tmp[1][j];
}
memset(tmp, 0, sizeof tmp);
for(register int i = 1; i <= in[len]; i++)
{
int res = 0;
memset(tmp, 0, sizeof tmp);
for(register int j = BL(i); j <= len; j++)
{
Extend(seq[j], tmp[0], tmp[1], res);
if(j == BR(in[j])) st[i][in[j]] = res;
}
}
return memset(tmp, 0, sizeof tmp), void();
}
inline i32 Run()
{
n = Read32(), m = Read32();
for(register int i = 1; i <= n; i++) seq[++len] = col[i] = Read32();
Discretization();
for(register int i = 1; i < n; i++)
{
int u = Read32(), v = Read32();
Link(u, v), Link(v, u);
}
DFS1(1, 0), DFS2(1, 1), Init();
while(m--)
{
int u = Read32() ^ ans, v = Read32(), t = LCA(u, v);
if(u != t && v != t)
{
if(pos[u][1] > pos[v][0]) swap(u, v);
int pu = pos[u][1], pv = pos[v][0];
if(in[pu] == in[pv])
{
ans = 0;
Extend(t, tmp[0], tmp[1], ans);
for(register int i = pu; i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
Write32(ans);
tmp[0][t] = tmp[1][col[t]] = 0;
for(register int i = pu; i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
continue;
}
int lb = in[pu] + 1, rb = in[pv] - 1;
ans = st[lb][rb];
tmp[0][t] = cnt[0][rb][t] - cnt[0][lb - 1][t], tmp[1][col[t]] = cnt[1][rb][col[t]] - cnt[1][lb - 1][col[t]];
for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
Extend(t, tmp[0], tmp[1], ans);
for(register int i = pu; i <= BR(in[pu]); i++) Extend(seq[i], tmp[0], tmp[1], ans);
for(register int i = BL(in[pv]); i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
Write32(ans);
tmp[0][t] = tmp[1][col[t]] = 0;
for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
}
else
{
if(pos[u][0] > pos[v][0]) swap(u, v);
int pu = pos[u][0], pv = pos[v][0];
if(in[pu] == in[pv])
{
ans = 0;
for(register int i = pu; i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
Write32(ans);
for(register int i = pu; i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
continue;
}
int lb = in[pu] + 1, rb = in[pv] - 1;
ans = st[lb][rb];
for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
for(register int i = pu; i <= BR(in[pu]); i++) Extend(seq[i], tmp[0], tmp[1], ans);
for(register int i = BL(in[pv]); i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
Write32(ans);
for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
}
}
return 0;
}
#undef BL
#undef BR
}
i32 main() { return Program :: Run(); }