一道简单的树形 dp 题。

思路

由于题目让我们求在树的子图中最长的蜈蚣图的长度,所以容易想到要用树形 dp 解决。我们定义一个蜈蚣图的身体为其中间的链,腿为和中间链连接的两个点。对于一个身体上的点来说,腿与其有两种连接方式,一种是选取两个子节点作为腿,一种是选取一个子节点和父节点作为腿。设当前考虑的节点为 $u$,如果选择了 $fa_u$,那么蜈蚣的身体就不能再向 $fa_u$ 及其祖先转移。所以我们定义 $f_{u,0}$ 表示以 $u$ 为蜈蚣的身体上端点而且不选 $fa_u$ 最为腿的最大长度,而 $f_{u,1}$ 则是选 $fa_u$ 最为腿的最大长度。然而我们发现身体端点不一定所有身体节点的祖先,所以我们定义 $g_u$ 表示以 $u$ 为身体中点的最大长度。

下面我们考虑如何转移。

若 $v$ 是 $u$ 的儿子,我们定义最大的 $f_{v,0}$ 叫 $w_0$,次大的 $f_{v,0}$ 叫 $w_1$,其儿子个数为 $n_u$。那么有以下几种转移。 $$ f_{u,0}=\max(f_{u,0},w_0+1),n_u\geq3 $$

$$ f_{u,1}=\max(f_{u,1},w_0+1),n_u\geq2\land u\neq root $$

当然,如果 $u$ 作为蜈蚣的身体下端点的话,则有如下转移: $$ f_{u,0}=\max(f_{u,0},1),n_u\geq2 $$

$$ f_{u,1}=\max(f_{u,1},1),n_u\geq1\land u\neq root $$

那么拼接两个蜈蚣的情况如下: $$ g_u=\max(g_u,w_0+w_1+1),n_u\geq4 \lor (n_u\geq3\land u\neq root) $$ 最后的答案如下: $$ ans=\max_{u=1}^n\left(\max(f_{u,0},f_{u,1},g_u)\right) $$

时间复杂度 $O(n)$。

代码

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
int n;
struct Node {
int to, next;
} edge[400010];
int head[200010], tot;
void add(int u, int v) {
edge[++tot].to = v;
edge[tot].next = head[u], head[u] = tot;
}
int dp[200010][2], Ans;
void dfs(int u, int fa) {
vector<int> v0;
for (int i = head[u]; i; i = edge[i].next) {
int v = edge[i].to;
if (v == fa) continue;
dfs(v, u);
v0.push_back(dp[v][0]);
}
sort(v0.begin(), v0.end(), greater<int>());
if (v0.size() >= 3) dp[u][0] = max(dp[u][0], v0[0] + 1);
if (v0.size() >= 2 && u != 1) dp[u][1] = max(dp[u][1], v0[0] + 1);
if (v0.size() >= 2) dp[u][0] = max(dp[u][0], 1);
if (v0.size() >= 1 && u != 1) dp[u][1] = max(dp[u][1], 1);
Ans = max({Ans, dp[u][0], dp[u][1]});
if (v0.size() >= 4) Ans = max(Ans, v0[0] + v0[1] + 1);
if (v0.size() >= 3 && u != 1) Ans = max(Ans, v0[0] + v0[1] + 1);
}
void init() {
for (int i = 1; i <= n; i++) head[i] = 0;
for (int i = 1; i <= tot; i++) edge[i].to = edge[i].next = 0;
for (int i = 1; i <= n; i++) dp[i][0] = dp[i][1] = 0;
Ans = 0;
}
void solve() {
n = read();
init();
for (int i = 1; i < n; i++) {
int u = read(), v = read();
add(u, v), add(v, u);
}
dfs(1, -1);
cout << Ans << endl;
}