You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何优化基于Numba的迭代式数组连通区域标记函数性能?

问题:优化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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.12 21:47:41