Python连通分量标记算法实现错误排查与修复咨询
连通分量标记算法错误排查与修复
问题描述
实现连通分量标记算法后,运行结果存在错误,例如标记为"3"和"5"的区域未合并。以下是实现代码及测试用输入数组:
原实现代码
import matplotlib.pyplot as plt from matplotlib import colors import numpy as np def resolve_equivalent_labels_dict(input_dict, labels): # merge lists for key in reversed(input_dict.keys()): input_dict[key] = list(set(sorted(input_dict[key]))) for key2, lis2 in input_dict.items(): if key != key2 and key in lis2: input_dict[key2].extend(input_dict[key]) input_dict[key] = [] # Fix problem, where a label is part of multiple lists. for label_i in labels: key_list = [] for key, lis in input_dict.items(): if label_i in lis: key_list.append(key) if len(key_list) > 1: key_list.remove(min(key_list)) for key in key_list: input_dict[key].remove(label_i) for key in input_dict.keys(): if input_dict[key] == []: input_dict.pop(key) else: input_dict[key] = sorted(list(set(input_dict[key]))) return input_dict def CCL(arr): # Write labels offsets = [(0,-1),(-1,0),(-1,-1),(-1,1)] equivalent_labels = {} current_label = 1 labels = [current_label] for row_i in range(len(arr)): for col_i in range(len(arr[0])): if arr[row_i, col_i] == 1: # Read labels from neighbors vals_neighbors = [] for (row_off, col_off) in offsets: v = arr[row_i+row_off, col_i+col_off] if v != 0: vals_neighbors.append(v) # Write labels and fill equivalent_label dict if len(vals_neighbors) == 0: # no neighbor, new label current_label+=1 labels.append(current_label) equivalent_labels[current_label] = [] arr[row_i, col_i] = current_label elif len(vals_neighbors) == 1: # one neighbor arr[row_i, col_i] = vals_neighbors[0] elif len(vals_neighbors) > 1: # multiple neighbors vals_neighbors = list(set(vals_neighbors)) first = min(vals_neighbors) vals_neighbors.remove(first) arr[row_i, col_i] = first if first in equivalent_labels.keys(): equivalent_labels[first].extend(vals_neighbors) else: equivalent_labels[first] = vals_neighbors # Merge lists print(equivalent_labels) equivalent_labels = resolve_equivalent_labels_dict(equivalent_labels, labels) print("") print(equivalent_labels) # Rename labels for key in reversed(equivalent_labels.keys()): lis = list(np.unique(equivalent_labels[key])) for elem in lis: arr = np.where(arr == elem, key, arr) return arr # create data gridlength = 20 input_array = np.zeros((gridlength+2, gridlength+2)) input_array[1:gridlength+1, 1:1+gridlength] = np.random.randint(2, size=(gridlength,gridlength)) print(input_array) # Plot fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(20,10)) ax1.pcolormesh(input_array, cmap=colors.ListedColormap(["white", "red"])) output_array = CCL(input_array) output_array[output_array==0] == np.nan ax2.pcolormesh(output_array) ax1.set_xlim([0,gridlength+2]) ax2.set_xlim([0,gridlength+2]) ax1.set_ylim([0,gridlength+2]) ax2.set_ylim([0,gridlength+2]) for (j,i),label in np.ndenumerate(output_array): ax2.text(i,j,int(label),ha='center',va='center') plt.savefig("code/connectivity/ccl.png", dpi=300, bbox_inches="tight")
测试输入数组
[[0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0.] [0. 1. 1. 0. 1. 1. 0. 0. 0. 1. 0. 1. 0. 0. 0. 1. 1. 0. 1. 0. 0. 0.] [0. 0. 0. 0. 1. 1. 0. 1. 1. 0. 0. 1. 0. 1. 0. 0. 0. 1. 1. 0. 1. 0.] [0. 0. 1. 1. 0. 0. 1. 1. 0. 0. 1. 0. 1. 0. 1. 0. 0. 1. 0. 1. 1. 0.] [0. 1. 1. 0. 1. 0. 0. 1. 1. 0. 1. 0. 1. 0. 1. 1. 1. 0. 1. 1. 0. 0.] [0. 1. 1. 1. 1. 1. 1. 0. 1. 0. 0. 0. 0. 0. 1. 1. 0. 0. 1. 1. 1. 0.] [0. 0. 1. 1. 1. 0. 0. 1. 0. 0. 1. 1. 0. 0. 0. 0. 1. 1. 1. 0. 1. 0.] [0. 1. 0. 0. 1. 0. 0. 0. 1. 0. 1. 1. 0. 1. 0. 0. 1. 0. 0. 0. 0. 0.] [0. 0. 0. 0. 1. 0. 0. 1. 1. 1. 1. 0. 0. 1. 0. 0. 1. 0. 0. 1. 1. 0.] [0. 0. 0. 0. 1. 1. 1. 0. 0. 1. 0. 0. 1. 1. 1. 0. 0. 1. 1. 1. 1. 0.] [0. 0. 1. 0. 1. 1. 0. 0. 1. 1. 0. 0. 1. 0. 1. 1. 0. 1. 1. 1. 1. 0.] [0. 0. 1. 0. 0. 1. 1. 0. 0. 1. 1. 0. 0. 0. 0. 0. 0. 0. 1. 0. 0. 0.] [0. 1. 0. 1. 1. 0. 1. 0. 1. 1. 0. 1. 0. 0. 0. 0. 1. 1. 1. 0. 0. 0.] [0. 1. 1. 1. 1. 1. 0. 1. 1. 0. 0. 0. 1. 0. 0. 0. 0. 1. 1. 0. 0. 0.] [0. 0. 1. 0. 0. 0. 0. 1. 0. 0. 1. 1. 0. 1. 1. 0. 0. 0. 1. 1. 0. 0.] [0. 1. 0. 0. 0. 1. 1. 0. 0. 0. 0. 0. 0. 0. 1. 0. 1. 1. 0. 1. 1. 0.] [0. 1. 1. 1. 0. 1. 1. 1. 0. 0. 0. 1. 1. 0. 1. 0. 1. 1. 1. 1. 0. 0.] [0. 1. 1. 1. 1. 0. 0. 0. 0. 1. 1. 0. 1. 1. 1. 0. 0. 1. 0. 0. 0. 0.] [0. 0. 1. 0. 0. 1. 1. 0. 0. 1. 0. 0. 0. 1. 0. 0. 1. 0. 1. 0. 0. 0.] [0. 0. 1. 0. 0. 0. 1. 1. 1. 1. 0. 1. 1. 1. 0. 1. 0. 0. 0. 1. 1. 0.] [0. 0. 1. 0. 0. 0. 1. 0. 1. 1. 1. 0. 1. 1. 1. 0. 0. 1. 0. 1. 0. 0.] [0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0.]]
错误原因分析
- 等价标签合并逻辑缺陷:
resolve_equivalent_labels_dict函数无法处理等价标签的传递关系(如A等价于B、B等价于C时,无法将三者合并为同一组),导致部分等价区域未合并。 - 邻域访问未做边界检查:直接访问
arr[row_i+row_off, col_i+col_off],边缘像素会读取数组反向索引(如最后一行),引入错误的邻域标签,干扰等价关系建立。 - 标签重命名逻辑不彻底:仅将等价标签替换为字典key,但未处理key之间的等价关系,导致部分区域仍保持不同标记。
- 绘图代码赋值错误:
output_array[output_array==0] == np.nan是比较操作而非赋值,无法隐藏背景0值。
修复方案
使用**并查集(Union-Find)**管理等价标签,这是连通分量标记中处理等价关系的标准方案,能高效且正确地合并所有等价区域。
修复后的完整代码
import matplotlib.pyplot as plt from matplotlib import colors import numpy as np class UnionFind: def __init__(self): self.parent = {} def find(self, x): if x not in self.parent: self.parent[x] = x if self.parent[x] != x: self.parent[x] = self.find(self.parent[x]) return self.parent[x] def union(self, x, y): x_root = self.find(x) y_root = self.find(y) if x_root != y_root: self.parent[y_root] = x_root def CCL(arr): # 8连通邻域(左、上、左上、右上) offsets = [(0,-1), (-1,0), (-1,-1), (-1,1)] uf = UnionFind() current_label = 1 # 第一遍扫描:标记标签并记录等价关系 for row_i in range(len(arr)): for col_i in range(len(arr[0])): if arr[row_i, col_i] == 1: neighbor_labels = [] for (row_off, col_off) in offsets: ni, nj = row_i + row_off, col_i + col_off # 边界检查,避免读取无效索引 if 0 <= ni < len(arr) and 0 <= nj < len(arr[0]): v = arr[ni, nj] if v != 0: neighbor_labels.append(v) if not neighbor_labels: # 无邻域标签,分配新标签 arr[row_i, col_i] = current_label uf.find(current_label) # 初始化标签的父节点 current_label += 1 else: # 取最小的邻域标签作为当前标签 min_label = min(neighbor_labels) arr[row_i, col_i] = min_label # 将所有邻域标签与最小标签合并 for label in neighbor_labels: uf.union(min_label, label) # 第二遍扫描:将所有标签替换为其根标签 for row_i in range(len(arr)): for col_i in range(len(arr[0])): label = arr[row_i, col_i] if label != 0: arr[row_i, col_i] = uf.find(label) return arr # 测试代码 gridlength = 20 # 使用用户提供的输入数组 input_array = np.array([ [0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.,0.], [0.,1.,1.,0.,1.,1.,0.,0.,0.,1.,0.,1.,0.,0.,0.,1.,1.,0.,1.,0.,0.,0.], [0.,0.,0.,0.,1.,1.,0.,1.,1.,0.,0.,1.,0.,1.,0.,0.,0.,1.,1.,0.,1.,0.], [0.,0.,1.,1.,0.,0.,1.,1.,0.,0.,1.,0.,1.,0.,1.,0.,0.,1.,0.,1.,1.,0.], [0.,1.,1.,0.,1.,0.,0.,1.,1.,0.,1.,0.,1.,0.,1.,1.,1.,0.,1.,1.,0.,0.], [0.,1.,1.,1.,1.,1.,1.,0.,1.,0.,0.,0.,0.,0.,1.,1.,0.,0.,1.,1.,1.,0.], [0.,0.,1.,1.,1.,0.,0.,1.,0.,0.,1.,1.,0.,0.,0.,0.,1.,1.,1.,0.,1.,0.], [0.,1.,0.,0.,1.,0.,0.,0.,1.,0.,1.,1.,0.,1.,0.,0.,1.,0.,0.,0.,0.,0.], [0.,0.,0.,0.,1.,0.,0.,1.,1.,1.,1.,0.,0.,1.,0.,0.,1.,0.,0.,1.,1.,0.], [0.,0.,0.,0.,1.,1.,1.,0.,0.,1.,0.,0.,1.,1.,1.,0.,0.,1.,1.,1.,1.,0.], [0.,0.,1.,0.,1.,1.,0.,0.,1.,1.,0.,0.,1.,0.,1.,1.,0.,1.,1.,1.,1.,0.], [0.,0.,1.,0.,0.,1.,1.,0.,0.,1.,1.,0.,0.,0.,0.,0.,0.,0.,1.,0.,0.,0.], [0.,1.,0.,1.,1.,0.,1.,0.,1.,1.,0.,1.,0.,0.,0.,0.,1.,1.,1.,0.,0.,0.], [0.,1.,1.,1.,1.,1.,0.,1.,1.,0.,0.,0.,1.,0.,0.,0.,0.,1.,1.,0.,0.,0.], [0.,0.,1.,0.,0.,0.,0.,1.,0.,0.,1.,1.,0.,1.,1.,0.,0.,0.,1.,1.,0.,0.], [0.,1.,0.,0.,0.,1.,1.,0.,0.,0.,0.,0.,0.,0.,1.,0.,1.,1.,0.,1.,1.,0.], [
相关产品推荐
相关产品推荐

