树中点对路径最值差之和求解:代码错误及正确方法问询
树中所有顶点对路径最大最小值差之和的正确解法
原代码错误原因
你写的代码完全忽略了树的结构,它计算的是所有顶点对(i,j)(i≤j)的顶点数值差的绝对值之和,但题目要求的是i到j的路径上所有顶点的最大值与最小值之差的总和。只有当树中任意两个顶点的路径就是它们本身(比如树是单点或两个直接相连的点)时,原代码才会巧合得到正确结果,其他场景必然出错。
正确思路
题目要求的总和可以拆分为:
总结果 = 所有路径的最大值之和 - 所有路径的最小值之和
因为对于任意路径(i,j),D(i,j)=max(path(i,j)) - min(path(i,j)),对所有i≤j求和后,等价于最大值的总和减去最小值的总和。
我们可以用贡献法分别计算这两个总和:
- 对每个顶点的数值
A[k],计算它作为多少条路径的最大值,将A[k] * 数量累加得到所有路径的最大值总和。 - 同理,计算每个顶点的数值作为多少条路径的最小值,累加得到最小值总和。
这里用并查集+排序的方法高效计算贡献,核心逻辑:
- 计算最大值贡献时,按数值从小到大处理顶点,每次将当前顶点与已处理的邻居合并,合并时新增的路径数是两个连通块大小的乘积,这些路径的最大值都是当前顶点的数值。
- 计算最小值贡献时,按数值从大到小处理顶点,逻辑类似,新增路径的最小值是当前顶点的数值。
Java 实现代码
import java.util.*; public class TreePathMaxMinSum { static class UnionFind { int[] parent; int[] size; public UnionFind(int n) { parent = new int[n]; size = new int[n]; for (int i = 0; i < n; i++) { parent[i] = i; size[i] = 1; } } public int find(int x) { if (parent[x] != x) { parent[x] = find(parent[x]); } return parent[x]; } public long union(int x, int y) { int rootX = find(x); int rootY = find(y); if (rootX == rootY) { return 0; } if (size[rootX] < size[rootY]) { int temp = rootX; rootX = rootY; rootY = temp; } long product = (long) size[rootX] * size[rootY]; parent[rootY] = rootX; size[rootX] += size[rootY]; return product; } } public static long solve(int[] A, int[] U, int[] V) { int n = A.length; List<List<Integer>> adj = new ArrayList<>(); for (int i = 0; i < n; i++) { adj.add(new ArrayList<>()); } // 转换为0-based索引(假设输入U、V是1-based) for (int i = 0; i < U.length; i++) { int u = U[i] - 1; int v = V[i] - 1; adj.get(u).add(v); adj.get(v).add(u); } long sumMax = calculateSum(A, adj, true); long sumMin = calculateSum(A, adj, false); return sumMax - sumMin; } private static long calculateSum(int[] A, List<List<Integer>> adj, boolean isMax) { int n = A.length; Integer[] nodes = new Integer[n]; for (int i = 0; i < n; i++) { nodes[i] = i; } // 排序:求max则升序,求min则降序 Arrays.sort(nodes, (a, b) -> isMax ? Integer.compare(A[a], A[b]) : Integer.compare(A[b], A[a])); UnionFind uf = new UnionFind(n); boolean[] visited = new boolean[n]; long sum = 0; // 先加入所有单点路径的贡献(i=j的情况) for (int num : A) { sum += num; } for (int u : nodes) { visited[u] = true; for (int v : adj.get(u)) { if (visited[v]) { long cnt = uf.union(u, v); sum += (long) A[u] * cnt; } } } return sum; } public static void main(String[] args) { // 测试示例:树为1-2-3,A=[1,2,3] int[] A = {1,2,3}; int[] U = {1,2}; int[] V = {2,3}; System.out.println(solve(A, U, V)); // 输出4,符合预期 } }
代码说明
- 邻接表构建:将输入的边数组转换为树的邻接表,处理1-based到0-based的索引转换。
- 并查集:维护连通块的大小,合并时快速计算新增的路径数量。
- 总和计算:通过排序节点,按顺序合并已访问的邻居,累加当前节点作为最大值/最小值的贡献。
- 结果输出:最大值总和减去最小值总和即为最终答案。
内容的提问来源于stack exchange,提问作者Learner
相关产品推荐
相关产品推荐

