Python 3:优化布尔掩码数组的连续非零段边界提取算法,实现代码简化
Simplify Extraction of Object Base/Top from Boolean Mask Array
我已经有一套可行的算法来从布尔掩码数组中识别连续对象区域的底部(每个区域的起始索引)和顶部(每个区域的结束索引),但想简化代码实现来提升技术能力。
场景说明
给定布尔掩码数组 obj_level,示例:
obj_level = [0, 0, 1, 1, 0, 0, 1, 1, 1, 1, 1, 0, 0, 0, 0, 1, 1]
需要提取每个连续1区域的第一个索引(base)和最后一个索引(top)。
原可行算法代码
import numpy as np base = [] top = [] obj_idx = np.flatnonzero(obj_level) if obj_idx.size > 0: base.append(obj_idx[0]) for i, idx in enumerate(obj_idx[:-1]): if idx + 1 == obj_idx[i+1]: continue else: top.append(idx) base.append(obj_idx[i+1]) top.append(obj_idx[-1])
我尝试简化成以下代码,但得到了布尔值和整数混合的数组(比如[True, True, 3, True, True]),希望找到比从混合数组提取整数更简单的实现方式:
base = [ idx + 1 == obj_idx[i+1] or idx+1 for i,idx in enumerate(obj_idx[:-1]) ] top = [ (idx+1 == obj_idx[i+1] or idx) for i,idx in enumerate(obj_idx[:-1]) ] np.insert(base,0,obj_idx[0]) np.insert(top,-1,obj_idx[-1])
更简洁的实现方案(基于Numpy)
你的核心需求是找到连续索引块的分割点,利用Numpy的diff函数可以高效定位这些分割点,完全避免循环和混合类型数组的问题。
实现代码:
import numpy as np obj_level = [0, 0, 1, 1, 0, 0, 1, 1, 1, 1, 1, 0, 0, 0, 0, 1, 1] obj_idx = np.flatnonzero(obj_level) if obj_idx.size > 0: # 计算相邻索引的差值,找到非连续的位置 splits = np.where(np.diff(obj_idx) > 1)[0] # 提取每个块的起始索引(base) base = obj_idx[np.concatenate([[0], splits + 1])] # 提取每个块的结束索引(top) top = obj_idx[np.concatenate([splits, [-1]])] else: base = np.array([]) top = np.array([]) print("Base:", base) # 输出: Base: [ 2 6 15] print("Top:", top) # 输出: Top: [ 3 10 16]
原理说明:
np.diff(obj_idx)计算相邻非零索引的差值,连续索引的差值为1,非连续的差值大于1;np.where(np.diff(obj_idx) > 1)[0]找到所有分割点的位置(即前一个块的最后一个元素在obj_idx中的索引);- 对于
base,我们需要每个块的第一个元素:第一个块从索引0开始,后续块从分割点的下一个位置(splits +1)开始; - 对于
top,我们需要每个块的最后一个元素:前几个块的最后元素就是分割点位置的元素,最后一个块的最后元素是obj_idx的最后一个元素(用[-1]表示)。
纯Python实现(无需Numpy)
如果你想不用Numpy,也可以用itertools.groupby来分组连续的索引,代码同样简洁:
from itertools import groupby obj_level = [0, 0, 1, 1, 0, 0, 1, 1, 1, 1, 1, 0, 0, 0, 0, 1, 1] # 先获取所有值为1的索引 obj_idx = [i for i, val in enumerate(obj_level) if val == 1] base = [] top = [] # 按连续索引分组 for _, group in groupby(obj_idx, key=lambda x: x - obj_idx.index(x)): group_list = list(group) base.append(group_list[0]) top.append(group_list[-1]) print("Base:", base) # 输出: Base: [2, 6, 15] print("Top:", top) # 输出: Top: [3, 10, 16]
原理说明:
groupby的key=lambda x: x - obj_idx.index(x)利用了连续索引的特性:连续的索引减去它们在obj_idx中的位置,结果是一个固定值,以此来分组连续块。
内容的提问来源于stack exchange,提问作者EJSABOLK
相关产品推荐
相关产品推荐

