【题解】THUWC 2018 城市规划

题目

思路

我们称 \(col_i\) 表示第 \(i\) 个点的颜色。

首先我们考虑相邻两个颜色都不相同怎么做,我们固定一个根,然后计算每个点当作连通块根的答案。

然后我们发现在特殊性质下,我们就只需要计算颜色交替的连通块数量就可以了,于是我们可以设 \(pre_i\) 表示在只能使用点 \(i\) 的颜色和父亲颜色的,且钦定点 \(i\) 要选的方案数,转移就是:

\[\normalsize{pre_u = \prod_{u \to v \land col_v = col_{fa}} (pre_v + 1)} \]

这个很好理解啊,感觉就是能从颜色和父亲颜色相同的点转移而来,同时乘上不选的可能性。

然后答案是什么呢?

我们现在计算以 \(u\) 为根的答案,我们发现我们需要枚举与 \(u\) 相邻的颜色 \(x\),然后答案就是:

\[\normalsize{ans_u = (\sum_x \prod_{u \to v \land col_v = x} pre_v + 1) - n + 1} \]

最后的减 \(n\) 是减去全不选跟自己颜色相同的点的,然后加一是算上只选自己的方案。

那么现在相邻点颜色有相同的,怎么办?

我们稍微修改一下 \(pre_i\) 的定义,定义为可使用的颜色为自己的颜色和最近的一个祖先和自己不同的颜色

然后我们思考一下答案应该变成什么形式:

我们发现一维不是很好描述这个东西,于是我们设 \(dp_{i,j}\) 表示以 \(i\) 为根的子树,可使用的第二种颜色 \(j\) 的方案数:

对于一条边 \(u \to v\),如果 \(col_u = col_v\),有:

\[dp_{u,j} = dp_{u,j} \times (dp_{v,j} + 1) \]

如果 \(col_u \neq col_v\),有:

\[dp_{u,col_v} = dp_{u,col_v} \times (pre_v + 1) \]

答案就是:

\[(\sum_{x=1}^n dp_{u,x}) - dp_{u, col_u} \times (n - 1) \]

然后我们发现,只要在 \(x\) 位置上不进行第二种转移,那么一定有 \(dp_{u,x} = dp_{u,col_u}\),于是我们发现就可以维护那些进行过第二种转移位置的值,这种值一共只有 \(O(n)\) 个,然后相当于我们需要维护的操作有全局加 \(1\),然后和某个数组相乘,这个是可以启发式合并的。

具体的说:我们对每个点维护有值的 dp 位置,总和,有值位置的数量,加法的懒标记,乘法的懒标记与逆元,以及单点相乘的函数。

然后启发式合并的时候我们就先全局称上被合并掉的那个点自己颜色的 dp 值,然后枚举被合并掉的 dp 位置,然后单点乘就可以了。

总的复杂度为 \(O(n \log n)\),瓶颈在于启发式合并。

代码

//奇跡を信じて,願いを叶えて
#include<bits/stdc++.h>
using namespace std;
#define int long long
#define double long double
#define uint unsigned long long
#define Air
namespace io{
	inline int read(){
		int f = 1, t = 0; char ch = getchar();
		while(ch < '0' || ch > '9'){if(ch == '-') f = -f; ch = getchar();}
		while(ch >= '0' && ch <= '9'){t = t * 10 + ch - '0'; ch = getchar();}
		return t * f;
	}
	inline void write(int x){
		if(x < 0){putchar('-'); x = -x;}
		if(x >= 10){write(x / 10);}
		putchar(x % 10 + '0');
	}
}
using namespace io;
int n;
const int N = 5e5 + 10, MOD = 998244353;
int a[N];
vector<int>e[N];
int lnk[N];
int pre[N];
int com[N];
int quick_pow(int a, int b){
	if(!b) return 1;
	if(b & 1) return a * quick_pow(a, b - 1) % MOD;
	else {
		int tmp = quick_pow(a, b / 2); return tmp * tmp % MOD;
	}
}
void get_lnk(int now, int ff){
	if(a[now] != a[ff]){
		lnk[now] = ff;
 	}
 	else{
 		lnk[now] = lnk[ff];
 	}
 	pre[now] = 1;
 	com[now] = 1;
	for(auto y: e[now]){
		if(y == ff) continue;
		get_lnk(y, now);
		if(a[y] == a[now] || a[y] == a[lnk[now]]){
			pre[now] *= (pre[y] + 1);
			pre[now] %= MOD;
		}
		if(a[y] == a[now]){
			com[now] *= (com[y] + 1);
		}
	}
}

int ans;

struct Node{
	unordered_map<int, int> dp;
	int laz, sum;
	int siz;
	int mul, inv;
	int col;
	void add_laz(int v){
 		laz += v;
 		sum += siz * v;
 		sum %= MOD;
	}
	void add_mul(int v){
 		mul *= v;
 		inv *= quick_pow(v, MOD - 2);
 		mul %= MOD;
 		inv %= MOD;
 		laz *= v;
 		laz %= MOD;
 		sum *= v;
 		sum %= MOD;
	}
	void poi_mul(int pos, int val){
		if(dp.find(pos) != dp.end()){
			sum -= (dp[pos] * mul % MOD + laz) % MOD;
			sum += MOD;
			sum += (dp[pos] * mul % MOD + laz) * val % MOD;
			sum %= MOD;
			dp[pos] = ((dp[pos] * mul % MOD + laz) * val - laz + MOD) % MOD * inv % MOD;
		}
		else{
			siz ++;
			sum += (dp[col] * mul % MOD + laz) * val % MOD;
			sum %= MOD;
			dp[pos] = ((dp[col] * mul % MOD + laz) * val % MOD - laz + MOD) % MOD * inv % MOD;
		}
	}
}val[N];
vector<int> bin;
int id[N];
int tot;
void dfs(int now, int ff){
	if(!bin.size()){
		val[++tot].dp.clear();
		val[tot].laz = val[tot].sum = val[tot].siz = 0;
		val[tot].mul = val[tot].inv = 1;
		val[tot].col = a[now];
		id[now] = tot;
	}
	else{
		val[bin.back()].dp.clear();
		val[bin.back()].laz = val[bin.back()].sum = val[bin.back()].siz = 0;
		val[bin.back()].col = a[now];
		val[bin.back()].mul = val[bin.back()].inv = 1;
		id[now] = bin.back();
		bin.pop_back();
	}
	val[id[now]].dp[a[now]] = 1;
	val[id[now]].siz = 1;
	val[id[now]].sum = 1;
	for(auto y: e[now]){
		if(y == ff) continue;
		dfs(y, now);
	}
	for(auto y: e[now]){
		if(y == ff) continue;
		if(a[y] == a[now]){
			val[id[y]].add_laz(1);
			if(val[id[y]].siz <= val[id[now]].siz){
				int v = val[id[y]].dp[a[now]] * val[id[y]].mul + val[id[y]].laz;
				v %= MOD;
				val[id[now]].add_mul(v);
				int in = quick_pow(v, MOD - 2);
				for(auto tt: val[id[y]].dp){
					val[id[now]].poi_mul(tt.first, (tt.second * val[id[y]].mul % MOD + val[id[y]].laz) % MOD * in % MOD);
				}
				bin.push_back(id[y]);
			}
			else{
				int v = val[id[now]].dp[a[now]] * val[id[now]].mul + val[id[now]].laz;
				v %= MOD;
				val[id[y]].add_mul(v);
				int in = quick_pow(v, MOD - 2);
				for(auto tt: val[id[now]].dp){
					val[id[y]].poi_mul(tt.first, (tt.second * val[id[now]].mul % MOD + val[id[now]].laz) % MOD * in % MOD);
				}
				bin.push_back(id[now]);
				id[now] = id[y];
			}
		}
		else{
			val[id[now]].poi_mul(a[y], pre[y] + 1);
		}
	}
	int sum = val[id[now]].sum - (val[id[now]].dp[a[now]] * val[id[now]].mul % MOD + val[id[now]].laz) % MOD * (val[id[now]].siz - 1) % MOD + MOD;
	sum = (sum % MOD + MOD) % MOD;
	sum %= MOD;
	ans += sum; ans %= MOD;
}
signed main() {
#ifndef Air
	freopen(".in","r",stdin);
	freopen(".out","w",stdout);
#endif
	ios::sync_with_stdio(false);cin.tie(0);cout.tie(0);
	n = read();
	for(int i = 1; i <= n; i++){
		a[i] = read();
	}
	for(int i = 1; i < n; i++){
		int x = read(), y = read();
		e[x].push_back(y);
		e[y].push_back(x);
	}
	get_lnk(1, 0);
	dfs(1, 0);
	cout << ans << '\n';
	return 0;
}
posted @ 2026-09-21 15:20  Air2011  阅读(8)  评论(0)    收藏  举报