CF2219C Coloring a Red Black Tree 题解
思路
树形 DP。设 \(dp_{u, 0/1}\) 代表把以 \(u\) 为根的子树全部染红,且 \(u\) 的父亲是黑/红时的期望操作次数。分两种情况讨论:
1. \(u\) 为红色:
此时不需要进一步对 \(u\) 染色,只需要将每一个子树染色即可:
\[dp_{u, 0/1} = \sum_{v \in \operatorname{son}(u)} dp_{v, 1}
\]
2. \(u\) 为黑色
我们会发现,将这个子树全部染红可以分为三个步骤:
- 选择 \(u\) 的一些子节点,将这些子节点及其子树全部染红;
- 将 \(u\) 染红;
- 将 \(u\) 剩下的子节点的子树全部染红。
那实际上就是枚举 \(\operatorname{son}(u)\) 的所有子集 \(S\),再求:
\[dp_{u, 0/1} = \min_S(\sum_{v \in S} dp_{v, 0} + val + \sum_{v \in \operatorname{son}(u), v \notin S} dp_{v, 1})
\]
其中 \(val\) 为在这种情况下将 \(u\) 染红的期望次数。我们会发现,如果每次向 \(S\) 中加入一个子节点,实际上的贡献是 \(dp_{v, 0} - dp_{v, 1}\)。考虑贪心,初始 \(S = \varnothing\),每次按 \(dp_{v, 0} - dp_{v, 1}\) 从大到小加入 \(S\),再求 \(\min\)。
注意:概率为 \(P\) 的事件发生的期望为 \(\dfrac{1}{P}\)。
时间复杂度 \(\mathcal{O}(n \log n)\),\(\log\) 是因为要 sort。
代码细节见注释。
代码
#include <bits/stdc++.h>
using namespace std;
const int N = 2e5 + 5;
int t, n, tmp;
int deg[N], s[N];
double dp[N][2];
vector<int> tr[N];
bool cmp(int x, int y)
{
// 这里判断 x, y == tmp 是为了将 u 的父亲甩到数组的最后面
if (x == tmp) return 0;
if (y == tmp) return 1;
return dp[x][1] - dp[x][0] > dp[y][1] - dp[y][0];
}
void dfs(int u, int fa)
{
dp[u][0] = dp[u][1] = 0;
for (int v : tr[u])
{
if (v == fa) continue;
dfs(v, u);
// u 已经是红色
if (s[u] == 1)
{
dp[u][0] += dp[v][1];
dp[u][1] += dp[v][1];
}
}
// u 还是黑色
if (s[u] == 0)
{
tmp = fa;
sort(tr[u].begin(), tr[u].end(), cmp);
double tot0 = 0, tot1 = 0;
for (int v : tr[u])
{
if (v == fa) continue;
tot1 += dp[v][1];
}
// 初始化 dp[u][0] 的含义:开始没有任何一个相邻点被染红,所以不可能将 u 染红
// 初始化 dp[u][1] 的含义:先染 u 的概率为 1 / deg[u](只有 u 的父亲被染红),再将所有子节点染红
dp[u][0] = 1e18, dp[u][1] = tot1 + deg[u]; // tot1 为先染红的集合的期望,tot2 为后染红的集合的期望
for (int i = 1; i <= tr[u].size(); i++)
{
int v = tr[u][i - 1];
if (v == fa) continue;
tot0 += dp[v][0], tot1 -= dp[v][1]; // 动态维护 tot0, tot1
// 将 u 染红的概率为 i / deg[u](i 即为提前染红的子节点),期望为 deg[u] / i
dp[u][0] = min(dp[u][0], tot0 + tot1 + 1.0 * deg[u] / i);
// 将 u 染红的概率为 (i + 1) / deg[u](i + 1 即为提前染红的子节点和 u 的父节点),期望为 deg[u] / (i + 1)
dp[u][1] = min(dp[u][1], tot0 + tot1 + 1.0 * deg[u] / (i + 1));
}
}
}
void solve()
{
scanf("%d", &n);
for (int i = 1; i <= n; i++)
scanf("%1d", &s[i]);
for (int i = 1, u, v; i < n; i++)
{
scanf("%d%d", &u, &v);
tr[u].push_back(v);
tr[v].push_back(u);
deg[u]++, deg[v]++;
}
dfs(1, 0);
// 最后答案为 dp[1][0],因为根节点没有父亲
printf("%.12lf\n", dp[1][0]);
for (int i = 1; i <= n; i++)
{
deg[i] = 0;
tr[i].clear();
}
}
int main()
{
scanf("%d", &t);
while (t--) solve();
return 0;
}

浙公网安备 33010602011771号