0pts
查看原帖
0pts
448887
cancan123456楼主2023/3/1 22:06
#include <cstdio>
#include <vector>
using namespace std;
const int N = 100005;
vector < int > to[N];
int w[N], fa[N], size_[N], dep[N], son[N];
void dfs1(int u, int fa) {
	::fa[u] = fa;
	size_[u] = 1;
	dep[u] = dep[fa] + 1;
	for (int v : to[u]) {
		if (v != fa) {
			dfs1(v, u);
			size_[u] += size_[v];
			if (size_[v] > size_[son[u]]) {
				son[u] = v;
			}
		}
	}
}
int top[N], dfn[N], timer, a[N];
void dfs2(int u, int toplink) {
	timer++;
	dfn[u] = timer;
	a[dfn[u]] = w[u];
	top[u] = toplink;
	if (son[u] != 0) {
		dfs2(son[u], toplink);
		for (int v : to[u]) {
			if (v != fa[u] && v != son[u]) {
				dfs2(v, v);
			}
		}
	}
}
struct Message {
	int lcol, rcol, cnt;
	Message() {
		lcol = rcol = 0;
		cnt = 0;
	}
	void set(int color) {
		lcol = rcol = color;
		cnt = 1;
	}
};
Message operator + (const Message & left, const Message & right) {
	Message sum;
	sum.lcol = left.lcol;
	sum.rcol = right.rcol;
	sum.cnt = left.cnt + right.cnt;
	if (left.rcol == right.lcol) {
		sum.cnt--;
	}
	return sum;
}
Message inv(Message m) {
	Message ans;
	ans.lcol = m.rcol;
	ans.rcol = m.lcol;
	ans.cnt = m.cnt;
	return ans;
}
struct Node {
	int l, r, tag;
	Message m;
} node[4 * N];
void push_up(int p) {
	node[p].m = node[2 * p].m + node[2 * p + 1].m;
}
void push_tag(int p, int tag) {
	if (tag != -1) {
		node[p].tag = tag;
		node[p].m.set(tag);
	}
}
void push_down(int p) {
	push_tag(2 * p, node[p].tag);
	push_tag(2 * p + 1, node[p].tag);
	node[p].tag = -1;
}
void build(int p, int l, int r) {
	node[p].l = l;
	node[p].r = r;
	node[p].tag = -1;
	if (l == r) {
		node[p].m.set(a[l]);
	} else {
		int mid = (l + r) / 2;
		build(2 * p, l, mid);
		build(2 * p + 1, mid + 1, r);
		push_up(p);
	}
}
void modify(int p, int l, int r, int tag) {
	if (l <= node[p].l && node[p].r <= r) {
		push_tag(p, tag);
	} else {
		push_down(p);
		int mid = (node[p].l + node[p].r) / 2;
		if (l <= mid) {
			modify(2 * p, l, r, tag);
		}
		if (mid + 1 <= r) {
			modify(2 * p + 1, l, r, tag);
		}
		push_up(p);
	}
}
Message query(int p, int l, int r) {
	if (l <= node[p].l && node[p].r <= r) {
		return node[p].m;
	} else {
		push_down(p);
		int mid = (node[p].l + node[p].r) / 2;
		if (r <= mid) {
			return query(2 * p, l, r);
		} else if (mid + 1 <= l) {
			return query(2 * p + 1, l, r);
		} else {
			return query(2 * p, l, r) + query(2 * p + 1, l, r);
		}
	}
}
void modify(int u, int v, int color) {
	while (top[u] != top[v]) {
		if (dep[top[u]] < dep[top[v]]) {
			u ^= v ^= u ^= v;
		}
		modify(1, dfn[top[u]], dfn[u], color);
		u = fa[top[u]];
	}
	if (dep[u] > dep[v]) {
		u ^= v ^= u ^= v;
	}
	modify(1, dfn[u], dfn[v], color);
}
Message query(int u, int v) {
	Message mu, mv;
	while (top[u] != top[v]) {
		if (dep[top[u]] > dep[top[v]]) {
			mu = query(1, dfn[top[u]], dfn[u]) + mu;
			u = fa[top[u]];
		} else {
			mv = query(1, dfn[top[v]], dfn[v]) + mv;
			v = fa[top[v]];
		}
	}
	if (dep[u] < dep[v]) {
		mu = query(1, dfn[u], dfn[v]) + mu;
	} else {
		mv = inv(query(1, dfn[v], dfn[u])) + mv;
	}
	return inv(mu) + mv;
}
int main() {
	int n, m;
	scanf("%d %d", &n, &m);
	for (int i = 1; i <= n; i++) {
		scanf("%d", &w[i]);
	}
	for (int u, v, i = 1; i < n; i++) {
		scanf("%d %d", &u, &v);
		to[u].push_back(v);
		to[v].push_back(u);
	}
	dfs1(1, 0);
	dfs2(1, 1);
	build(1, 1, n);
	char op[2];
	for (int u, v, c, i = 1; i <= m; i++) {
		scanf("%s %d %d", op, &u, &v);
		if (op[0] == 'C') {
			scanf("%d", &c);
			modify(u, v, c);
		} else {
			printf("%d\n", query(u, v).cnt);
		}
	}
	return 0;
}
2023/3/1 22:06
加载中...