如何优化基于Numba的迭代式数组连通区域标记函数性能?
我正在使用Python结合Numba编写一个2D或3D数组的对象标记函数,目标是将输入数组中所有正交连通且值相同的单元,在输出数组中赋予从1到N的唯一标记(N为正交连通组的数量)。该功能与scipy.ndimage.label等库函数类似,但这类函数会将所有正交连通的非零单元归为同一标记,会合并不同值的连通组,不符合需求。
例如输入:
[0 0 7 7 0 0 0 0 7 0 0 0 0 0 0 0 0 7 0 6 6 0 0 7 0 0 4 4 0 0]
scipy.ndimage.label会将6和4合并为标记2,而我需要的输出是:
[0 0 1 1 0 0 0 0 1 0 0 0 0 0 0 0 0 4 0 2 2 0 0 4 0 0 3 3 0 0]
现有实现中,递归版经Numba JIT编译后仅需1-2秒,但处理大型数组时会触发递归深度限制导致崩溃;改写为迭代版后,Numba编译下耗时约10秒,仍有优化空间。
递归版代码:
@numba.njit def adjacent(idx, shape): coords = [] if len(shape) > 2: if idx[0] < shape[0] - 1: coords.append((idx[0] + 1, idx[1], idx[2])) if idx[0] > 0: coords.append((idx[0] - 1, idx[1], idx[2])) if idx[1] < shape[1] - 1: coords.append((idx[0], idx[1] + 1, idx[2])) if idx[1] > 0: coords.append((idx[0], idx[1] - 1, idx[2])) if idx[2] < shape[2] - 1: coords.append((idx[0], idx[1], idx[2] + 1)) if idx[2] > 0: coords.append((idx[0], idx[1], idx[2] - 1)) else: if idx[0] < shape[0] - 1: coords.append((idx[0] + 1, idx[1])) if idx[0] > 0: coords.append((idx[0] - 1, idx[1])) if idx[1] < shape[1] - 1: coords.append((idx[0], idx[1] + 1)) if idx[1] > 0: coords.append((idx[0], idx[1] - 1)) return coords @numba.njit def apply_label(labels, decoded_image, current_label, idx): labels[idx] = current_label for aidx in adjacent(idx, labels.shape): if decoded_image[aidx] == decoded_image[idx] and labels[aidx] == 0: apply_label(labels, decoded_image, current_label, aidx) @numba.njit def label_image(decoded_image): labels = np.zeros_like(decoded_image, dtype=np.uint32) current_label = 0 for idx in zip(*np.where(decoded_image >= 0)): if labels[idx] == 0: current_label += 1 apply_label(labels, decoded_image, current_label, idx) return labels, current_label
迭代版代码:
@numba.njit def label_image(decoded_image): labels = np.zeros_like(decoded_image, dtype=np.uint32) current_label = 0 for idx in zip(*np.where(decoded_image >= 0)): if labels[idx] == 0: current_label += 1 idxs = [idx] while idxs: cidx = idxs.pop() if labels[cidx] == 0: labels[cidx] = current_label for aidx in adjacent(cidx, labels.shape): if labels[aidx] == 0 and decoded_image[aidx] == decoded_image[idx]: idxs.append(aidx) return labels, current_label
请问如何优化该迭代版本的性能,使其接近递归版本的效率?
优化方案
1. 提取基准值,减少重复数组访问
迭代版中每次判断decoded_image[aidx] == decoded_image[idx]时,decoded_image[idx]是当前连通组的固定基准值,可提前提取存储,避免反复访问数组:
# 在处理当前连通组时添加 target_val = decoded_image[idx] # 后续判断改为 if labels[aidx] == 0 and decoded_image[aidx] == target_val:
2. 重构邻域生成逻辑,消除动态列表开销
原adjacent函数每次创建新列表,在Numba中会产生不必要的内存分配开销。改用预定义偏移量的方式直接计算邻域坐标,避免动态列表:
@numba.njit def get_neighbors(idx, shape): neighbors = [] # 预定义正交偏移量 if len(shape) == 2: offsets = [(-1, 0), (1, 0), (0, -1), (0, 1)] else: offsets = [(-1,0,0), (1,0,0), (0,-1,0), (0,1,0), (0,0,-1), (0,0,1)] for offset in offsets: if len(shape) == 2: ni, nj = idx[0] + offset[0], idx[1] + offset[1] if 0 <= ni < shape[0] and 0 <= nj < shape[1]: neighbors.append((ni, nj)) else: ni, nj, nk = idx[0]+offset[0], idx[1]+offset[1], idx[2]+offset[2] if 0 <= ni < shape[0] and 0 <= nj < shape[1] and 0 <= nk < shape[2]: neighbors.append((ni, nj, nk)) return neighbors
这种方式减少了条件分支的嵌套,同时避免了动态列表的频繁创建。
3. 提前标记已访问节点,避免重复入栈
原迭代版在弹出栈元素时才标记,可能导致同一坐标被多次加入栈。改为在将邻域坐标加入栈之前就标记,减少栈的操作次数:
# 处理当前连通组时 current_label += 1 target_val = decoded_image[idx] labels[idx] = current_label # 提前标记起始点 idxs = [idx] while idxs: cidx = idxs.pop() for aidx in get_neighbors(cidx, labels.shape): if labels[aidx] == 0 and decoded_image[aidx] == target_val: labels[aidx] = current_label # 入栈前标记 idxs.append(aidx)
这样可以避免同一个坐标被多次添加到栈中,减少冗余的栈操作和判断。
4. 使用Numba类型化列表替代普通Python列表
普通Python列表在Numba中的优化有限,改用numba.typed.List可以让Numba更好地优化栈的操作:
import numba # 在函数内初始化栈时 idxs = numba.typed.List() idxs.append(idx)
注意需要提前导入numba并确保类型化列表正确初始化。
5. 优化遍历方式,避免np.where的额外开销
原代码用zip(*np.where(decoded_image >= 0))遍历元素,会生成中间坐标数组,开销较大。改用直接嵌套循环遍历数组索引,在Numba中能获得更高的效率:
@numba.njit def label_image(decoded_image): labels = np.zeros_like(decoded_image, dtype=np.uint32) current_label = 0 shape = decoded_image.shape if len(shape) == 2: for i in range(shape[0]): for j in range(shape[1]): idx = (i, j) if labels[idx] == 0: current_label += 1 target_val = decoded_image[idx] labels[idx] = current_label idxs = numba.typed.List() idxs.append(idx) while idxs: cidx = idxs.pop() for aidx in get_neighbors(cidx, shape): if labels[aidx] == 0 and decoded_image[aidx] == target_val: labels[aidx] = current_label idxs.append(aidx) elif len(shape) == 3: for i in range(shape[0]): for j in range(shape[1]): for k in range(shape[2]): idx = (i, j, k) if labels[idx] == 0: current_label += 1 target_val = decoded_image[idx] labels[idx] = current_label idxs = numba.typed.List() idxs.append(idx) while idxs: cidx = idxs.pop() for aidx in get_neighbors(cidx, shape): if labels[aidx] == 0 and decoded_image[aidx] == target_val: labels[aidx] = current_label idxs.append(aidx) return labels, current_label
直接遍历索引避免了np.where生成中间数组的内存和时间开销,Numba对嵌套循环的优化也更充分。
内容的提问来源于stack exchange,提问作者Colin

