用的是拆 16 个询问的莫队。
#include <cstdio>
#include <algorithm>
using namespace std;
const int N = 100005;
const int M = 500005;
const int B = 300;
struct Edge {
int v, next;
} edge[2 * N];
int head[N];
int cnt;
void add_edge(int u, int v) {
cnt++;
edge[cnt].v = v;
edge[cnt].next = head[u];
head[u] = cnt;
}
int a[N], b[N];
int dfn[N], dfncnt;
int sub_s[N], sub_t[N];
int temp_a[N];
int f[N][17], dep[N];
void dfs(int u, int fa) {
dep[u] = dep[fa] + 1;
f[u][0] = fa;
for (int i = 1; i < 17; i++) {
f[u][i] = f[f[u][i - 1]][i - 1];
}
dfncnt++;
dfn[u] = dfncnt;
temp_a[dfn[u]] = a[u];
sub_s[u] = dfncnt;
for (int v, i = head[u]; i != 0; i = edge[i].next) {
v = edge[i].v;
if (v != fa) {
dfs(v, u);
}
}
sub_t[u] = dfncnt;
}
struct Query {
int x, y, id, sign;
} query[16 * M];
bool operator < (const Query & a, const Query & b) {
return a.x / B != b.x / B ? a.x / B < b.x / B : a.y < b.y;
}
long long ans_q[4 * M];
int query_cnt, id_cnt;
void __add_query(int x, int y, int id, int sign) {
if (x > y) {
swap(x, y);
}
query_cnt++;
query[query_cnt] = (Query) {x, y, id, sign};
}
int belong[4 * M];
int n;
void add_query(int l1, int r1, int l2, int r2, int iid) {
if (l1 <= 0 || l1 > n) {
return;
}
if (r1 <= 0 || r1 > n) {
return;
}
if (l2 <= 0 || l2 > n) {
return;
}
if (r2 <= 0 || r2 > n) {
return;
}
id_cnt++;
belong[id_cnt] = iid;
__add_query(r1, r2, id_cnt, 1);
__add_query(l2 - 1, r1, id_cnt, -1);
__add_query(l1 - 1, r2, id_cnt, -1);
__add_query(l1 - 1, l2 - 1, id_cnt, 1);
}
int cnt1[N], cnt2[N];
long long nowans;
void insert1(int x) {
nowans += cnt2[x];
cnt1[x]++;
}
void insert2(int x) {
nowans += cnt1[x];
cnt2[x]++;
}
void delete1(int x) {
nowans -= cnt2[x];
cnt1[x]--;
}
void delete2(int x) {
nowans -= cnt1[x];
cnt2[x]--;
}
long long ans[M];
int main() {
int m;
scanf("%d %d", &n, &m);
for (int i = 1; i <= n; i++) {
scanf("%d", a + i);
b[i] = a[i];
}
sort(b + 1, b + n + 1);
int n_ = unique(b + 1, b + n + 1) - (b + 1);
for (int i = 1; i <= n; i++) {
a[i] = lower_bound(b + 1, b + n_ + 1, a[i]) - b;
}
for (int x, y, i = 1; i < n; i++) {
scanf("%d %d", &x, &y);
add_edge(x, y);
add_edge(y, x);
}
dfs(1, 0);
int root = 1;
int lx[2], rx[2], cntx;
int ly[2], ry[2], cnty;
int operation1_cnt = 0;
for (int op, x, y, i = 1; i <= m; i++) {
scanf("%d", &op);
if (op == 1) {
scanf("%d", &x);
root = x;
operation1_cnt++;
} else {
scanf("%d %d", &x, &y);
if (x == root) {
cntx = 1;
lx[0] = 1;
rx[0] = n;
} else if (sub_s[x] <= dfn[root] && dfn[root] <= sub_t[x]) {
cntx = 2;
int p = root;
int diff = dep[root] - (dep[x] + 1);
for (int i = 0; i < 17; i++) {
if (((diff >> i) & 1) == 1) {
p = f[p][i];
}
}
lx[0] = 1;
rx[0] = sub_s[p] - 1;
lx[1] = sub_t[p] + 1;
rx[1] = n;
} else {
cntx = 1;
lx[0] = sub_s[x];
rx[0] = sub_t[x];
}
if (y == root) {
cnty = 1;
ly[0] = 1;
ry[0] = n;
} else if (sub_s[y] <= dfn[root] && dfn[root] <= sub_t[y]) {
cnty = 2;
int p = root;
int diff = dep[root] - (dep[y] + 1);
for (int i = 0; i < 17; i++) {
if (((diff >> i) & 1) == 1) {
p = f[p][i];
}
}
ly[0] = 1;
ry[0] = sub_s[p] - 1;
ly[1] = sub_t[p] + 1;
ry[1] = n;
} else {
cnty = 1;
ly[0] = sub_s[y];
ry[0] = sub_t[y];
}
for (int i0 = 0; i0 < cntx; i0++) {
for (int i1 = 0; i1 < cnty; i1++) {
add_query(lx[i0], rx[i0], ly[i1], ry[i1], i - operation1_cnt);
}
}
}
}
int l = 0, r = 0;
sort(query + 1, query + query_cnt + 1);
for (int x, y, i = 1; i <= query_cnt; i++) {
x = query[i].x;
y = query[i].y;
while (l < x) {
l++;
insert2(a[l]);
}
while (l > x) {
delete2(a[l]);
l--;
}
while (r < y) {
r++;
insert1(a[r]);
}
while (r > y) {
delete1(a[r]);
r--;
}
ans_q[query[i].id] += query[i].sign * nowans;
}
for (int i = 1; i <= id_cnt; i++) {
ans[belong[i]] += ans_q[i];
}
for (int i = 1; i <= m - operation1_cnt; i++) {
printf("%lld\n", ans[i]);
}
return 0;
}