CF2178F 思路分享(组合数学,并查集)

https://codeforces.com/problemset/problem/2178/F

题意概述

给定一棵根为 \(1\) 的树,若节点 \(i\) 子树大小为偶数,\(i\) 为白色,否则为黑色.

\(i\) 的父节点为 \(f\),对每个白色的节点,可以断开 \(i\)\(f\) 的边,再任意添加一条边,需要保证操作后图仍为树.

定义一棵树 \(conquer\):所有白色的节点都在以 \(1\) 为起点的某条链上.

求给定树经过任意次操作能得到的 \(conquer\) 树的数量,模 \(998244353\).

\(2\le n \le 2\cdot 10^5\).

思路

先断掉所有白色节点与父亲的边,树被分成若干连通块,记为 \(S_0,S_1,\cdots,S_k\)\(S_0\)\(1\) 所属的连通块.

需要把连通块串成链,\(S_0\) 为头.

任意两连通块 \(S_i,S_j\) 连接方案数为 \(|S_i| \cdot |S_j|\).

对于链 \(S_0,S_{p_1},S_{p_2},\cdots,S_{p_k}\),方案数为

\[|S_0|\cdot |S_{p_1}|^2 \cdot |S_{p_2}|^2 \cdots |S_{p_k}| \]

枚举链尾连通块,前面 \(k-1\) 个连通块任意排列,总贡献为

\[|S_0| \cdot (k-1)! \cdot \prod_{i=1}^{k}{|S_i|^2} \cdot \sum_{i=1}^{k}{\frac{1}{|S_i|}} \]

时间复杂度 \(\mathcal{O}(n)\).

代码

//author:kzssCCC

#include <bits/stdc++.h>
using namespace std;
using ll = long long;

class dsu{
public:
	int n,cnt_cc;
	vector<int> p,sz;
	
	dsu(int _n){
		n = _n;
		cnt_cc = n;
		p = vector<int>(n+1);
		for (int i=1;i<=n;i++){
			p[i] = i;
		}
		
		sz = vector<int>(n+1,1);
	}
	
	int find(int x){
		int root = x;
		while (p[root]!=root) root = p[root];
		
		while (x!=root){
			int next = p[x];
			p[x] = root;
			x = next;
		}
		
		return root;
	}
	
	void unite(int a,int b){
		a = find(a);
		b = find(b);
		
		if (a==b) return;
		
		if (sz[a]>=sz[b]){
			sz[a] += sz[b];
			p[b] = a;
		}
		else{
			sz[b] += sz[a];
			p[a] = b;
		}
		
		cnt_cc--;
	}
};

const int MOD = 998244353;

ll qpow(ll a,ll b){
	ll res = 1;
	while (b){
		if (b&1){
			res = res*a%MOD;
		}
		a = a*a%MOD;
		b >>= 1;
	}

	return res;
}

void solve(){
	int n;
	cin >> n;

	vector<vector<int>> adj(n+1);
	for (int i=0;i<n-1;i++){
		int u,v;
		cin >> u >> v;
		adj[u].push_back(v);
		adj[v].push_back(u);
	}

	dsu ds(n);
	vector<int> sz(n+1);
	function<void(int,int)> dfs = [&](int u,int par){
		sz[u] = 1;
		for (auto& v:adj[u]){
			if (v==par) continue;
			dfs(v,u);
			sz[u] += sz[v];

			if (sz[v]&1){
				ds.unite(u,v);
			}
		}
	};

	dfs(1,-1);
	map<int,int> mp;
	int rt;
	for (int i=1;i<=n;i++){
		if (i==1){
			rt = ds.find(i);
		}

		mp[ds.find(i)]++;
	}

	if (mp.size()==1){
		cout << 1 << '\n';
		return;
	}

	ll res = mp[rt];
	ll sum = 0;
	for (auto& [v,c]:mp){
		if (v==rt) continue;
		sum = (sum+qpow(c,MOD-2))%MOD;
		res = res*c%MOD*c%MOD;
	}

	res = res*sum%MOD;
	int k = mp.size()-1;
	for (int i=2;i<k;i++){
		res = res*i%MOD;
	}
	
	cout << res << '\n';
}

int main(){
	ios::sync_with_stdio(false);
	cin.tie(0);
	
	int t = 1;
	cin >> t;
	while (t--) solve();

	return 0;
}
posted @ 2026-06-16 17:28  kzssCCC  阅读(6)  评论(0)    收藏  举报