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;
}
posted @ 2026-07-02 23:47  kzssCCC  阅读(4)  评论(0)    收藏  举报