CF2219C 思路分享(dp,期望,并查集)
https://codeforces.com/problemset/problem/2219/C
题意
给定 \(n\) 个节点的无根树和字符串 \(s\),\(s_i=1\) 表示初始 \(i\) 节点为红色,\(s_i=0\) 表示初始为黑色.
每次操作选择一个节点 \(u\),随机选择一个 \(u\) 的邻居 \(v\),将 \(u\) 染成 \(v\) 的颜色.
求让所有节点为红色的期望最少操作次数.
\(1\le n\le 2\cdot 10^5\).
思路
将红色节点视为间隔,将树划分成若干棵子树,任意选一个黑色节点作为根,红色节点只可能作为叶子,独立处理所有子树.
发现贪心不可行,考虑 \(dp\).
处理节点 \(u\) 时,除了子节点,还需要知道其父节点的颜色状态,这可以通过染色的先后顺序刻画.
定义 \(dp[u][0]\) 为,\(u\) 先于其父节点染成红色,处理完 \(u\) 及子树时的贡献;\(dp[u][1]\) 为,\(u\) 后于其父节点染成红色的贡献.
对于红色节点 \(u\),初始化 \(dp[u][0]=0\),\(dp[u][1] =+ \infty\).
转移过程中,记 \(cnt\) 为 \(u\) 邻居中本来就是红色的节点数量,这类节点不参与讨论,将所有 \(u\) 的子树中取红色的节点归于 \(S\) 集,取黑色的节点归于 \(T\) 集.
\[dp[u][0] = \min(\frac{deg_u}{cnt+|S|}+\sum_{v\in S}{dp[v][0]}+\sum_{v\in T}{dp[v][1]})
\]
可以枚举 \(|S|\),令 \(cur = \sum_{v\in S}{dp[v][0]}+\sum_{v\in T}{dp[v][1]}\),本质就是最小化 \(cur\).
初始全部归于 \(T\) 集,\(cur=\sum{dp[v][1]}\),\(|S|\) 加 \(1\) 时,贪心地取最小的 \(dp[v][0]-dp[v][1]\) 即可.
时间复杂度 \(\mathcal{O}(n\log 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 double INF = 1e18;
const double eps = 1e-9;
int sgn(double x){
if (abs(x)<=eps) return 0;
else if (x>0) return 1;
else return -1;
}
void solve(){
int n;
string s;
cin >> n >> s;
s = ' '+s;
dsu ds(n);
vector<vector<int>> adj(n+1);
vector<int> deg(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);
if (s[u]=='0' && s[v]=='0'){
ds.unite(u,v);
}
deg[u]++,deg[v]++;
}
set<int> st;
for (int i=1;i<=n;i++){
st.insert(ds.find(i));
}
vector<array<double,2>> dp(n+1,array<double,2>{INF,INF});
function<void(int,int)> dfs = [&](int u,int par){
if (s[u]=='1'){
dp[u][0] = 0.0;
dp[u][1] = INF;
return;
}
int cnt = 0;
for (auto& v:adj[u]){
if (v==par) continue;
dfs(v,u);
if (s[v]=='1'){
cnt++;
}
}
sort(adj[u].begin(),adj[u].end(),[&](int i,int j){
return sgn((dp[i][0]-dp[i][1])-(dp[j][0]-dp[j][1]))<0;
});
double cur = 0;
for (auto& v:adj[u]){
if (v==par || s[v]=='1') continue;
cur += dp[v][1];
}
int f = 0;
for (int k=cnt;1;k++){
if (k>0 && sgn(cur+1.0*deg[u]/k-dp[u][0])<0){
dp[u][0] = cur+1.0*deg[u]/k;
}
if (sgn(cur+1.0*deg[u]/(k+1)-dp[u][1])<0){
dp[u][1] = cur+1.0*deg[u]/(k+1);
}
while (f<deg[u] && (adj[u][f]==par || s[adj[u][f]]=='1')){
f++;
}
if (f<deg[u]){
int v = adj[u][f++];
cur += dp[v][0]-dp[v][1];
}
else break;
}
};
double res = 0;
for (auto& rt:st){
if (s[rt]=='1') continue;
dfs(rt,-1);
res += dp[rt][0];
}
cout << fixed << setprecision(12) << res << '\n';
}
int main(){
ios::sync_with_stdio(false);
cin.tie(0);
int t = 1;
cin >> t;
while (t--) solve();
return 0;
}

浙公网安备 33010602011771号