如何基于一维二进制掩码提取连续1分块子数组(PyTorch/NumPy)
需求说明
给定一维二进制掩码数组,例如np.array([0,0,0,1,1,1,0,0,1,1,0]),需要提取数组中所有连续为1的区域,拆分为独立子数组块。
原有演示代码拆分后会同时包含全0块和全1块:
import torch import numpy as np arr = torch.from_numpy(np.array([0,0,0,1,1,1,0,0,1,1,0])) split_idx = torch.where(torch.diff(arr) == 1)[0]+1 torch.tensor_split(arr, split_idx.tolist())
运行输出为:
(tensor([0, 0, 0]), tensor([1, 1, 1]), tensor([0, 0]), tensor([1, 1]), tensor([0]))
期望输出仅保留全1的连续块:
(tensor([1, 1, 1]), tensor([1, 1]))
实现方案
PyTorch 实现
简洁版(基于原有逻辑修改)
原有拆分逻辑已经将连续同值区域切分完成,只需新增过滤步骤保留全1块即可:
split_idx = torch.where(torch.diff(arr) == 1)[0] + 1 splits = torch.tensor_split(arr, split_idx.tolist()) # 过滤全1块 one_blocks = tuple(block for block in splits if torch.all(block == 1))
高效版(直接定位1块索引)
针对超长数组场景,可直接定位每个连续1块的起止位置,无需拆分全量数组再过滤:
# 数组首尾补0,统一处理边界处的1块 pad_arr = torch.nn.functional.pad(arr, (1, 1), value=0) # 提取0→1的起始位置、1→0的结束位置 starts = torch.where(torch.diff(pad_arr) == 1)[0] ends = torch.where(torch.diff(pad_arr) == -1)[0] one_blocks = tuple(arr[s:e] for s, e in zip(starts, ends))
两种写法输出完全一致,均符合预期。
NumPy 实现
逻辑和PyTorch版本完全对应:
arr = np.array([0,0,0,1,1,1,0,0,1,1,0]) # 简洁版 split_idx = np.where(np.diff(arr) == 1)[0] + 1 splits = np.array_split(arr, split_idx) one_blocks = tuple(block for block in splits if np.all(block == 1)) # 高效版 pad_arr = np.pad(arr, (1,1), constant_values=0) starts = np.where(np.diff(pad_arr) == 1)[0] ends = np.where(np.diff(pad_arr) == -1)[0] one_blocks = tuple(arr[s:e] for s,e in zip(starts, ends))
内容的提问来源于stack exchange,提问作者JXuan
相关产品推荐
相关产品推荐

