Kruskal算法含指定边实现问题:最小生成树关键边识别错误排查
排查LeetCode「找出最小生成树中的关键边和伪关键边」代码错误
问题描述
我正在解决LeetCode上的《找出最小生成树中的关键边和伪关键边》问题:
- 关键边:所有最小生成树都必须包含的边
- 伪关键边:存在至少一个最小生成树包含的边
我实现了基于并查集的Kruskal算法,以及强制包含指定边的Kruskal变体,但在测试用例:n=6,edges=[[0,1,1],[1,2,1],[0,2,1],[2,3,4],[3,4,2],[3,5,2],[4,5,2]]
时,输出关键边为空、伪关键边为[0,1,2,3,4,5,6],而预期结果是关键边[3]、伪关键边[0,1,2,4,5,6],请求帮忙排查代码错误。
代码实现
class UnionFind: def __init__(self, n): self.parent = [i for i in range(n)] self.size = [1 for _ in range(n)] self.cnt = n def find(self, node): while node != self.parent[node]: self.parent[node] = self.parent[self.parent[node]] node = self.parent[node] return node def union_(self, node1, node2): root1 = self.find(node1) root2 = self.find(node2) if root1 == root2: return if self.size[root1] > self.size[root2]: self.parent[root2] = root1 self.size[root1] += 1 else: self.parent[root1] = root2 self.size[root2] += 1 self.cnt -= 1 class Solution: def kruskal(self, num_nodes, edges): result = 0 edges.sort(key = lambda x: x[2]) uf = UnionFind(num_nodes) for a, b, w in edges: if uf.find(a) != uf.find(b): uf.union_(a, b) result += w return (result, uf.cnt) def kruskal_included(self, num_nodes, edges, edge): result = edge[2] edges.sort(key = lambda x: x[2]) uf = UnionFind(num_nodes) uf.union_(edge[0], edge[1]) for a, b, w in edges: if uf.find(a) != uf.find(b): uf.union_(a, b) result += w return result def findCriticalAndPseudoCriticalEdges(self, n: int, edges: List[List[int]]) -> List[List[int]]: min_cost, _ = self.kruskal(n, edges) first_list = [] second_list = [] for i in range(len(edges)): new_edges = [edges[j] for j in range(len(edges)) if j != i] new_cost, cnt = self.kruskal(n, new_edges) if cnt > 1 or new_cost > min_cost: first_list.append(i) elif self.kruskal_included(n, edges, edges[i]) == min_cost: second_list.append(i) return [first_list, second_list]
错误分析与修复方案
1. 原地排序修改原数组,导致边索引混乱
这是测试用例错误的核心原因:
kruskal和kruskal_included方法中调用edges.sort()会原地修改传入的edges列表。- 第一次调用
kruskal获取最小代价后,原edges数组已被排序,后续循环中使用的edges不再是输入的原始顺序,导致i对应的边完全偏离题目原始索引,无法正确识别关键边(比如原始索引3的边被排到最后,循环到对应位置时i已不是3)。
修复方法:
排序前复制数组,避免修改原数组:
def kruskal(self, num_nodes, edges): result = 0 # 复制数组后排序,不修改原数组 sorted_edges = sorted(edges, key=lambda x: x[2]) uf = UnionFind(num_nodes) for a, b, w in sorted_edges: if uf.find(a) != uf.find(b): uf.union_(a, b) result += w return (result, uf.cnt)
def kruskal_included(self, num_nodes, edges, edge): result = edge[2] # 复制数组后排序 sorted_edges = sorted(edges, key=lambda x: x[2]) uf = UnionFind(num_nodes) uf.union_(edge[0], edge[1]) for a, b, w in sorted_edges: # 跳过强制包含的边,避免重复处理 if [a, b, w] == edge: continue if uf.find(a) != uf.find(b): uf.union_(a, b) result += w return result
2. 并查集Union方法的size更新错误
UnionFind的union_方法中,合并集合时size更新逻辑错误:
- 原代码用
self.size[root1] += 1仅给集合大小加1,正确做法是将被合并集合的完整大小加到目标集合上,否则会导致按秩合并逻辑失效。
修复方法:
def union_(self, node1, node2): root1 = self.find(node1) root2 = self.find(node2) if root1 == root2: return if self.size[root1] > self.size[root2]: self.parent[root2] = root1 self.size[root1] += self.size[root2] # 加上对方集合的完整大小 else: self.parent[root1] = root2 self.size[root2] += self.size[root1] # 加上对方集合的完整大小 self.cnt -= 1
3. 伪关键边判断逻辑冗余(可选优化)
题目要求关键边和伪关键边互斥,原代码会将关键边也加入伪关键边列表,需调整判断顺序:
for i in range(len(edges)): is_critical = False new_edges = [edges[j] for j in range(len(edges)) if j != i] new_cost, cnt = self.kruskal(n, new_edges) if cnt > 1 or new_cost > min_cost: first_list.append(i) is_critical = True # 非关键边才判断是否为伪关键边 if not is_critical and self.kruskal_included(n, edges, edges[i]) == min_cost: second_list.append(i)
修复后测试结果
修复上述问题后,针对目标测试用例会得到正确输出:关键边[3],伪关键边[0,1,2,4,5,6]。
内容的提问来源于stack exchange,提问作者S10000
相关产品推荐
相关产品推荐

