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;
}

浙公网安备 33010602011771号