C语言实现DSU时出现无明显原因的段错误(疑似低级错误)
我用C语言实现了带权重总和的Disjoint Set Union(DSU),程序包含5个节点,通过union_by_rank函数合并集合,find_sum函数用于查询指定节点所在集合的权重总和。例如节点权重为{11,13,1,3,5},当节点1-2-3相连、0-4相连时,节点2所在集合的总和应为13+1+3=17。
但运行时发现,find_sum函数中调用find_parent会触发段错误,而union_by_rank中调用find_parent却能正常工作。即使在find_sum中直接打印部分节点的sum值(如dsu->parent[3].sum为17、dsu->parent[0].sum为11)都正常,请问这是什么原因?
原代码
#include <stdio.h> #include <stdlib.h> typedef struct Parent { int node; int sum; } Parent; typedef struct DSU { Parent* parent; int* rank; } DSU; void create_dsu(DSU* dsu, int n, int* wts) { dsu->parent = malloc(sizeof(Parent) * n); dsu->rank = malloc(sizeof(int) * n); for (int i = 0; i < n; i++) { dsu->parent[i].sum = wts[i]; dsu->parent[i].node = i; dsu->rank[i] = 0; } } int find_parent(DSU* dsu, int n) { if (n == dsu->parent[n].node) { return n; } return dsu->parent[n].node = find_parent(dsu, n); } void union_by_rank(DSU* dsu, int u, int v) { int up = find_parent(dsu, u); int vp = find_parent(dsu, v); if (up == vp) { return; } else if (dsu->rank[up] > dsu->rank[vp]) { dsu->parent[vp].node = up; dsu->parent[up].sum += dsu->parent[vp].sum; } else if (dsu->rank[vp] > dsu->rank[up]) { dsu->parent[up].node = vp; dsu->parent[vp].sum += dsu->parent[up].sum; } else { dsu->parent[up].node = vp; dsu->rank[vp]++; dsu->parent[vp].sum += dsu->parent[up].sum; } } int find_sum(DSU* dsu, int u) { int up = find_parent(dsu, u); // causes a segfault // printf("%d\n", dsu->parent[3].sum); -> 17 // printf("%d\n", dsu->parent[0].sum); -> 11 return (dsu->parent[up].sum); } int main() { int arr[] = { 11, 13, 1, 3, 5 }; DSU dsu; create_dsu(&dsu, 5, arr); union_by_rank(&dsu, 1, 3); union_by_rank(&dsu, 2, 3); union_by_rank(&dsu, 0, 4); printf("%d\n", find_sum(&dsu, 2)); }
问题根源:find_parent函数的递归参数错误
看find_parent的实现:
int find_parent(DSU* dsu, int n) { if (n == dsu->parent[n].node) { return n; } return dsu->parent[n].node = find_parent(dsu, n); }
当节点n的父节点不是自己时,递归调用的参数仍然是n,而非当前节点的父节点dsu->parent[n].node。这会导致无限递归,最终栈溢出触发段错误。
为什么union_by_rank中调用没报错?因为在你的测试流程里,union_by_rank中的find_parent调用大多是处理父节点为自身的节点,没有触发递归;即使有递归,次数也极少,还没达到栈溢出的阈值。而find_sum中查询节点2时,节点2的父节点是3,此时调用find_parent(2)会进入递归,但参数还是2,导致无限循环调用find_parent(dsu,2),最终栈溢出。
修复方案
修改find_parent函数,递归时传入当前节点的父节点,同时完成路径压缩:
int find_parent(DSU* dsu, int n) { if (n == dsu->parent[n].node) { return n; } return dsu->parent[n].node = find_parent(dsu, dsu->parent[n].node); }
修复后,调用find_sum(&dsu,2)会正确返回17,程序正常运行。
另外,建议在程序结束时释放malloc分配的内存,避免内存泄漏:
int main() { int arr[] = { 11, 13, 1, 3, 5 }; DSU dsu; create_dsu(&dsu, 5, arr); union_by_rank(&dsu, 1, 3); union_by_rank(&dsu, 2, 3); union_by_rank(&dsu, 0, 4); printf("%d\n", find_sum(&dsu, 2)); // 释放内存 free(dsu.parent); free(dsu.rank); return 0; }
内容的提问来源于stack exchange,提问作者Aryan Kadole

