求助 WA 30 pts
查看原帖
求助 WA 30 pts
448887
cancan123456楼主2022/4/21 12:11

用的是拆 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;
}
2022/4/21 12:11
加载中...