P3177 树上染色 - 洛谷

题意

有一棵 $n$ 个节点的树,树边有边权。你要从中选择 $m$ 个点染成黑色,其余 $n-m$ 个点染成白色。

求在所有的染色方案中,黑点两两之间的距离和 + 白点两两之间的距离和 的最大值。

思路

如果直接去算“黑点两两距离 + 白点两两距离”,由于要枚举点对,时间复杂度直接炸掉TTTLLLEEE。

用点不行,我们试试用边:考虑每一条边被经过了多少次

一、 核心思路:算每条边的贡献

对于树上的任意一条边 $(u, v)$(假设 $v$ 是 $u$ 的子节点),如果我们将这条边断开,树会被分成两部分:

  1. $v$ 的子树(假设大小为 $siz[v]$,其中有 $k$ 个黑点,则有 $siz[v] - k$ 个白点)。

  2. 树的其余部分(共有 $n - siz[v]$ 个点,其中有 $m - k$ 个黑点,其余 $(n - siz[v]) - (m - k)$ 个为白点)。

这条边会被哪些点对跨越并经过呢?

  • 黑点对:子树内的 $k$ 个黑点 $\times$ 子树外的 $(m - k)$ 个黑点。

  • 白点对:子树内的 $(siz[v] - k)$ 个白点 $\times$ 子树外的 $(n - siz[v] - m + k)$ 个白点。

也就是说,只要确定了 $v$ 的子树内选了多少个黑点 $k$,这条边 $w$ 对总收益的贡献就是固定的:

$$\text{贡献} = w \times \Big( k \times (m - k) + (siz[v] - k) \times (n - siz[v] - m + k) \Big)$$

这样就把庞大的全局距离和,拆解成了每条边根据“子树黑点数”计算的局部贡献!

二、 树形 DP 状态设计与转移

依据上面的推导,我们需要维护子树内选了多少个黑点。

1. 状态定义

定义 $f[u][j]$ 表示:在以 $u$ 为根的子树中,选择 $j$ 个点染成黑色,该子树内部的所有边所产生的最大贡献和。

2. 状态转移(树上背包)

用树形背包的套路,遍历 $u$ 的每一个子节点 $v$:

  • 设当前 $u$ 已经处理过的子树大小为 $siz[u]$,$v$ 的子树大小为 $siz[v]$。

  • 状态转移核心思想就是:合并前 $u$ 选了 $j-k$ 个黑点,子树 $v$ 选了 $k$ 个黑点,两者合并后总共是 $j$ 个黑点。

分块代码

看完思路后,我们一步一步来分块拆解代码。

1. DFS 遍历与边权贡献计算

在合并子节点前,我们先递归处理子树,并计算当前边产生的跨越贡献。

1
2
3
4
5
6
7
void dfs(int u, int p){
siz[u] = 1;
for(auto [v, w] : adj[u]){
if(v == p) continue;
dfs(v, u); // 递归处理子树
// 注意:这里提前把 siz[v] 加到了 siz[u] 里面
siz[u] += siz[v];

2. 背包边界优化

j 表示合并后 u 子树的总黑点数。倒序枚举防止一个物品被多次放入(经典01背包)
k 表示从 v 子树中选出的黑点数。
下界 $max(0, j - siz[u] + siz[v])$ 是怎么来的?
此时的 $siz[u]$ 已经加上了 $siz[v]$,所以合并前 $u$ 的真实大小是 $(siz[u] - siz[v])$。我们从 $u$ 原本的状态转移来,$u$ 原本选了 $(j - k)$ 个黑点,那么必然有:$(j - k) \le siz[u] - siz[v] \implies k \ge j - siz[u] + siz[v]$
这样就防止新状态覆盖旧状态导致重复计算

1
2
3
4
5
6
7
8
9
        for(int j = min(m, siz[u]); j >= 0; --j){
for(int k = max(0, j - siz[u] + siz[v]); k <= min(siz[v], j); ++k){
ll sum = k * (m - k) + (siz[v] - k) * (n - siz[v] - m + k);
// 状态转移:取原最优解 vs (合并前u选j-k个) + (v选k个) + (当前边跨越贡献)
f[u][j] = max(f[u][j], f[u][j - k] + f[v][k] + w * sum);
}
}
}
}

最终代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int maxn=2000+5;
ll f[maxn][maxn];
struct Edge{ int u; ll w;};
vector<Edge> adj[maxn];
int siz[maxn];
int n,m;
void dfs(int u,int p){
siz[u]=1;
for(auto [v,w]:adj[u]){
if(v==p) continue;
dfs(v,u);
siz[u]+=siz[v];
for(int j=min(m,siz[u]);j>=0;--j){
for(int k=max(0,j-siz[u]+siz[v]);k<=min(siz[v],j);++k){
ll sum=k*(m-k)+(siz[v]-k)*(n-siz[v]-m+k);
f[u][j]=max(f[u][j],f[u][j-k]+f[v][k]+w*sum);
}
}
}
}
void solve()
{
cin>>n>>m;
for(int i=1;i<n;++i){
int u,v,w;
cin>>u>>v>>w;
adj[u].push_back({v,w});
adj[v].push_back({u,w});
}
dfs(1,0);
cout<<f[1][m]<<"\n";
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int ___T=1;
//cin>>___T;
while(___T--) solve();
return 0;
}

关于树上背包的时间复杂度

很多人看到嵌套的 dfs 加上里面两层 for 循环,会直觉性地认为这三个维度的枚举会导致时间复杂度达到 $O(N \times M^2)$ 甚至是 $O(N^3)$。

但实际上,只要加上了 min(m, siz[u])min(siz[v], j) 这两个极其精准的上下界,整个算法的复杂度是严格的 $O(N \times M)$

每次背包的合并,本质上是把 $u$ 原先子树里的节点 和 $v$ 子树里的节点两两组合。树上的任意两个节点,只会在它们的 LCA处被合并计算一次。 因此总的操作(合并)次数等于总的节点对数。受到背包容量 $m$ 的截断限制后,最终均摊下来就是 $O(N \times M)$。对于 $N, M \le 2000$ 的数据规模,能轻松跑完。