每日算法 — 使用java实现树形动态规划:子树状态递推与换根技巧

树形动态规划(Tree DP)是将动态规划思想应用于树结构的一类经典算法。与线性DP不同,树形DP利用树的递归结构,自底向上地在每个子树上计算最优解,最终合并到根节点得到全局答案。本文将从三个经典模型入手,用Java完整实现树形DP的核心代码,帮助读者建立树结构上的DP思维。

一、树形DP的核心思想

1.1 为什么需要树形DP

线性结构上的DP通常从左到右递推,但树结构没有天然的线性顺序。树形DP的突破口在于:树可以递归分解为若干子树。对于任意节点,以它为根的子树问题可以分解为以它的孩子为根的子树问题之和,再加上该节点本身的决策。

树形DP的一般框架如下:

  • 状态定义dp[u] 表示以节点 u 为根的子树在某条件下的最优解。
  • 状态转移:遍历 u 的所有孩子 v,将 dp[v] 合并到 dp[u] 中。
  • 递归边界:叶子节点的 dp 值直接确定。
  • 遍历顺序:后序遍历(先处理孩子,再处理自己)。

1.2 树的存储结构

树形DP通常使用邻接表存储树结构。为了避免递归时回到父节点,需要在DFS时传入父节点参数。

import java.util.*;

/**
 * 树的邻接表存储与基础遍历框架
 */
class Tree {
    int n; // 节点数
    List<Integer>[] adj; // 邻接表

    @SuppressWarnings("unchecked")
    Tree(int n) {
        this.n = n;
        adj = new ArrayList[n];
        for (int i = 0; i < n; i++) {
            adj[i] = new ArrayList<>();
        }
    }

    void addEdge(int u, int v) {
        adj[u].add(v);
        adj[v].add(u); // 无向树
    }

    /**
     * 树形DP基础DFS框架
     * @param u 当前节点
     * @param parent 父节点,防止回退
     */
    void dfs(int u, int parent) {
        for (int v : adj[u]) {
            if (v == parent) continue; // 不回到父节点
            dfs(v, u); // 先递归处理子树
            // 在这里将dp[v]的结果合并到dp[u]
        }
        // 处理u自身的DP状态
    }
}

二、经典模型一:树的最大独立集

2.1 问题描述

给定一棵树,每个节点有权值(或仅计数),要求选出一个节点子集,使得子集中任意两个节点都不相邻(没有边直接相连),且权值之和最大。

2.2 状态设计

对于每个节点 u,定义两个状态:

  • dp[u][0]:不选节点 u 时,以 u 为根的子树能获得的最大权值。
  • dp[u][1]:选节点 u 时,以 u 为根的子树能获得的最大权值。

2.3 状态转移

  • 不选u:孩子可以选也可以不选,取最大值之和。
  • dp[u][0] = Σ max(dp[v][0], dp[v][1])
  • 选u:孩子一定不能选。
  • dp[u][1] = weight[u] + Σ dp[v][0]

2.4 Java实现

/**
 * 树的最大独立集
 * 每个节点有权重,求不相邻节点的最大权值和
 */
class TreeMaxIndependentSet {
    int n;
    List<Integer>[] adj;
    int[] weight;
    // dp[u][0]: 不选u; dp[u][1]: 选u
    int[][] dp;

    @SuppressWarnings("unchecked")
    TreeMaxIndependentSet(int n, int[] weight) {
        this.n = n;
        this.weight = weight;
        this.dp = new int[n][2];
        adj = new ArrayList[n];
        for (int i = 0; i < n; i++) adj[i] = new ArrayList<>();
    }

    void addEdge(int u, int v) {
        adj[u].add(v);
        adj[v].add(u);
    }

    void solve(int root) {
        dfs(root, -1);
        int ans = Math.max(dp[root][0], dp[root][1]);
        System.out.println("最大独立集权值和: " + ans);
    }

    private void dfs(int u, int parent) {
        // 初始化:不选u时贡献为0,选u时贡献为自身权重
        dp[u][0] = 0;
        dp[u][1] = weight[u];

        for (int v : adj[u]) {
            if (v == parent) continue;
            dfs(v, u);
            // 不选u:孩子可选可不选
            dp[u][0] += Math.max(dp[v][0], dp[v][1]);
            // 选u:孩子必须不选
            dp[u][1] += dp[v][0];
        }
    }

    // 输出选了哪些节点(根据dp值回溯)
    List<Integer> getSelectedNodes(int root) {
        List<Integer> selected = new ArrayList<>();
        traceback(root, -1, -1, selected);
        return selected;
    }

    private void traceback(int u, int parent, int state, List<Integer> selected) {
        // state = -1 表示根节点,根据dp值决定
        if (state == -1) {
            state = dp[u][1] > dp[u][0] ? 1 : 0;
        }
        if (state == 1) {
            selected.add(u);
        }
        for (int v : adj[u]) {
            if (v == parent) continue;
            if (state == 1) {
                // 选了u,孩子必须不选
                traceback(v, u, 0, selected);
            } else {
                // 没选u,孩子根据哪个更大决定
                int childState = dp[v][1] > dp[v][0] ? 1 : 0;
                traceback(v, u, childState, selected);
            }
        }
    }
}

三、经典模型二:树的直径

3.1 问题描述

树的直径是树中任意两节点之间最长路径的边数(或权值和)。求树的直径是树形DP的经典应用之一。

3.2 两种思路

树形DP求直径的核心思想是:对于每个节点 u,经过 u 的最长路径等于 u 到各子树的最长向下路径中,取最长的两条相加。

定义 down[u] 为从 u 出发向子树方向能到达的最远距离。则:

  • down[u] = max(down[v] + w(u,v)),其中 vu 的孩子。
  • 经过 u 的最长路径候选值 = 最大的两个 down[v] + w(u,v) 之和。
  • 全局直径 = 所有节点候选值的最大值。

3.3 Java实现

/**
 * 树的直径(边权版)
 * 使用树形DP一次DFS求解
 */
class TreeDiameter {
    int n;
    List<int[]>[] adj; // 存储 [邻居, 边权]
    int diameter;
    int[] down; // down[u]: 从u向下最远的路径长度

    @SuppressWarnings("unchecked")
    TreeDiameter(int n) {
        this.n = n;
        this.diameter = 0;
        this.down = new int[n];
        adj = new ArrayList[n];
        for (int i = 0; i < n; i++) adj[i] = new ArrayList<>();
    }

    void addEdge(int u, int v, int w) {
        adj[u].add(new int[]{v, w});
        adj[v].add(new int[]{u, w});
    }

    int solve(int root) {
        dfs(root, -1);
        return diameter;
    }

    private void dfs(int u, int parent) {
        // 收集u到所有子树的向下最远距离
        int firstMax = 0;  // 最大的down[v] + w
        int secondMax = 0; // 第二大的

        for (int[] edge : adj[u]) {
            int v = edge[0];
            int w = edge[1];
            if (v == parent) continue;
            dfs(v, u);
            int dist = down[v] + w;
            if (dist > firstMax) {
                secondMax = firstMax;
                firstMax = dist;
            } else if (dist > secondMax) {
                secondMax = dist;
            }
        }

        down[u] = firstMax;
        // 经过u的最长路径 = 最长分支 + 次长分支
        diameter = Math.max(diameter, firstMax + secondMax);
    }
}

3.4 无权树简化版

如果所有边权为1,代码可以进一步简化:

/**
 * 无权树的直径(边数)
 */
class TreeDiameterUnweighted {
    int n;
    List<Integer>[] adj;
    int diameter;

    @SuppressWarnings("unchecked")
    TreeDiameterUnweighted(int n) {
        this.n = n;
        adj = new ArrayList[n];
        for (int i = 0; i < n; i++) adj[i] = new ArrayList<>();
    }

    void addEdge(int u, int v) {
        adj[u].add(v);
        adj[v].add(u);
    }

    int solve(int root) {
        dfs(root, -1);
        return diameter;
    }

    private int dfs(int u, int parent) {
        int firstMax = 0, secondMax = 0;
        for (int v : adj[u]) {
            if (v == parent) continue;
            int depth = dfs(v, u) + 1;
            if (depth > firstMax) {
                secondMax = firstMax;
                firstMax = depth;
            } else if (depth > secondMax) {
                secondMax = depth;
            }
        }
        diameter = Math.max(diameter, firstMax + secondMax);
        return firstMax;
    }
}

四、经典模型三:换根DP

4.1 问题描述

换根DP(Rerooting Technique)解决的是这样一类问题:对于树中的每个节点,如果以它为根,某个DP值是多少?暴力对每个节点做一次DFS是O(n²),换根DP通过两次DFS在O(n)时间内解决。

经典例题:求每个节点作为根时,树的高度(即该节点到最远叶子的距离)。

4.2 核心思想

换根DP分为两步:

  1. 第一次DFS(自底向上):计算每个节点 u 在其子树内的最远距离 down[u]
  2. 第二次DFS(自顶向下):计算每个节点 u 通过父节点方向能到达的最远距离 up[u]

对于节点 u,以它为根时的高度 = max(down[u], up[u])

关键在于如何计算 up[u]

  • up[u] 来自父节点 p 的 “通过非u分支能到达的最远距离” + w(p,u)。
  • 如果 down[u] + w(p,u) 恰好是 p 的最大分支,则需要用 p 的次大分支。

4.3 Java实现

/**
 * 换根DP:求以每个节点为根时的树高度
 * 即每个节点到最远叶子的距离
 */
class RerootDP {
    int n;
    List<int[]>[] adj;
    int[] down;  // down[u]: u到子树内最远叶子的距离
    int[] up;    // up[u]: u通过父节点方向能到达的最远距离
    int[] height; // height[u]: 以u为根时的树高度

    @SuppressWarnings("unchecked")
    RerootDP(int n) {
        this.n = n;
        down = new int[n];
        up = new int[n];
        height = new int[n];
        adj = new ArrayList[n];
        for (int i = 0; i < n; i++) adj[i] = new ArrayList<>();
    }

    void addEdge(int u, int v, int w) {
        adj[u].add(new int[]{v, w});
        adj[v].add(new int[]{u, w});
    }

    void solve(int root) {
        dfs1(root, -1); // 第一次DFS:计算down数组
        dfs2(root, -1); // 第二次DFS:计算up数组
        for (int i = 0; i < n; i++) {
            height[i] = Math.max(down[i], up[i]);
        }
    }

    // 第一次DFS:自底向上求down
    private void dfs1(int u, int parent) {
        for (int[] e : adj[u]) {
            int v = e[0], w = e[1];
            if (v == parent) continue;
            dfs1(v, u);
            down[u] = Math.max(down[u], down[v] + w);
        }
    }

    // 第二次DFS:自顶向下求up
    private void dfs2(int u, int parent) {
        // 先收集u的所有子节点的down值
        int m = adj[u].size();
        int[] childDown = new int[m];
        int[] childId = new int[m];
        int idx = 0;
        for (int[] e : adj[u]) {
            int v = e[0], w = e[1];
            childId[idx] = v;
            if (v == parent) {
                childDown[idx] = -1; // 标记父节点方向
            } else {
                childDown[idx] = down[v] + w;
            }
            idx++;
        }

        // 预处理前缀最大值和后缀最大值
        // 这样对于每个子节点v,可以快速获取"u的其他分支中的最大值"
        int[] prefixMax = new int[m];
        int[] suffixMax = new int[m];
        for (int i = 0; i < m; i++) {
            prefixMax[i] = (i == 0) ? childDown[i] : Math.max(prefixMax[i-1], childDown[i]);
        }
        for (int i = m - 1; i >= 0; i--) {
            suffixMax[i] = (i == m - 1) ? childDown[i] : Math.max(suffixMax[i+1], childDown[i]);
        }

        // 计算每个子节点的up值
        for (int i = 0; i < m; i++) {
            int v = childId[i];
            if (v == parent) continue;
            int w = 0;
            for (int[] e : adj[u]) if (e[0] == v) { w = e[1]; break; }

            // u的其他分支的最大值
            int otherMax = 0;
            if (m > 1) {
                int left = (i > 0) ? prefixMax[i-1] : -1;
                int right = (i < m - 1) ? suffixMax[i+1] : -1;
                otherMax = Math.max(left, right);
                if (otherMax < 0) otherMax = 0;
            }

            // up[v] = max(通过u的其他分支, up[u]) + w(u,v)
            up[v] = Math.max(otherMax, up[u]) + w;
            dfs2(v, u);
        }
    }
}

4.4 换根DP的通用模板

换根DP的关键是:每个节点的答案 = 子树内信息 + 子树外信息。通过前缀/后缀最大值技巧,可以在O(deg(u))时间内为每个子节点计算排除该子节点后的最优值。

/**
 * 换根DP通用模板框架
 * 适用于"以每个节点为根时的某种聚合值"问题
 */
abstract class RerootTemplate {
    int n;
    List<Integer>[] adj;

    // 子树内的DP值
    abstract long[] dfs1(int u, int parent);

    // 结合子树外信息,计算每个节点的最终答案
    abstract void dfs2(int u, int parent, long fromParent);

    // 合并多个子节点的结果(用于前缀/后缀处理)
    abstract long merge(long a, long b);
}

五、完整测试与运行示例

public class TreeDPDemo {
    public static void main(String[] args) {
        System.out.println("=== 测试1:树的最大独立集 ===");
        int n1 = 6;
        int[] weight = {10, 20, 30, 40, 50, 60};
        TreeMaxIndependentSet mis = new TreeMaxIndependentSet(n1, weight);
        mis.addEdge(0, 1);
        mis.addEdge(0, 2);
        mis.addEdge(1, 3);
        mis.addEdge(1, 4);
        mis.addEdge(2, 5);
        mis.solve(0);
        System.out.println("选中节点: " + mis.getSelectedNodes(0));
        // 最优解:选0,3,4,5 -> 10+40+50+60 = 160
        // 或选1,2 -> 20+30 = 50(不是最优)
        // 实际上 dp[0][0] = max(dp[1])+max(dp[2])...

        System.out.println("\n=== 测试2:树的直径 ===");
        TreeDiameter td = new TreeDiameter(6);
        td.addEdge(0, 1, 3);
        td.addEdge(0, 2, 1);
        td.addEdge(1, 3, 2);
        td.addEdge(1, 4, 4);
        td.addEdge(2, 5, 5);
        System.out.println("树的直径: " + td.solve(0)); // 3->1->0->2->5 = 3+3+1+5 = 12? 不对,应该是 4->1->0->2->5 = 4+3+1+5 = 13

        System.out.println("\n=== 测试3:换根DP ===");
        RerootDP rr = new RerootDP(6);
        rr.addEdge(0, 1, 3);
        rr.addEdge(0, 2, 1);
        rr.addEdge(1, 3, 2);
        rr.addEdge(1, 4, 4);
        rr.addEdge(2, 5, 5);
        rr.solve(0);
        System.out.println("以各节点为根时的高度:");
        for (int i = 0; i < 6; i++) {
            System.out.printf("  节点%d: height=%d%n", i, rr.height[i]);
        }
    }
}

六、复杂度分析

算法模型 时间复杂度 空间复杂度 说明
最大独立集 O(n) O(n) 每个节点访问一次,常数状态数
树的直径 O(n) O(n) 一次DFS,维护两个最大值
换根DP O(n) O(n) 两次DFS,前缀/后缀预处理

树形DP的时间复杂度通常为 O(n),因为每个节点只被访问常数次。空间复杂度也是 O(n),主要用于存储邻接表和DP数组。

七、总结与扩展

本文通过三个经典模型展示了树形DP的核心技巧:

  1. 状态设计与转移:利用树的后序遍历性质,自底向上合并子树信息。最大独立集展示了”选/不选”二类状态设计。
  2. 全局最优的局部维护:树的直径展示了如何在每个节点维护多个最优值(最大/次大分支),从而推导全局最优。
  3. 换根技巧:通过两次DFS结合前缀/后缀最大值,将O(n²)的暴力枚举优化到O(n)。

可扩展方向
树形背包:每个子树视为一个物品组,在容量限制下做分组背包,时间复杂度O(n×V)。
点分治:处理树上路径问题,每次找到重心分治,将问题规模减半。
虚树:在大量查询中只保留关键节点构建虚树,大幅降低树形DP的复杂度。