如何用Numpy获取二进制数组的子数组索引?
用Numpy优化连续1的区间索引提取
问题描述
给定数组:
r = np.array([1, 0, 0, 1, 1, 1, 0, 1, 1, 1])
期望输出所有连续1的起始和结束索引:
[(0, 0), (3, 5), (7, 9)]
当前使用Python循环实现的函数如下,但处理200万元素的数组时耗时约8秒,速度较慢:
def get_indicies(array): indicies = [] xstart = None for x, col in enumerate(array): if col == 0 and xstart is not None: indicies.append((xstart, x - 1)) xstart = None elif col == 1 and xstart is None: xstart = x if xstart is not None: indicies.append((xstart, x)) return indicies
希望通过Numpy内置方法(如argwhere、split等)提升处理速度。
Numpy优化方案
利用Numpy的向量化操作替代Python循环,能大幅提升处理效率,具体实现如下:
import numpy as np def get_continuous_ones(arr): # 给数组前后补0,处理开头/结尾为连续1的边界情况 padded = np.concatenate(([0], arr, [0])) # 计算差分,定位0和1的转换点 diff = np.diff(padded) # 提取所有连续1的起始索引(0→1的上升沿) starts = np.where(diff == 1)[0] # 提取所有连续1的结束索引(1→0的下降沿,需减1修正) ends = np.where(diff == -1)[0] - 1 # 配对起始和结束索引为元组列表 return list(zip(starts, ends))
测试验证
用示例数组测试:
r = np.array([1, 0, 0, 1, 1, 1, 0, 1, 1, 1]) print(get_continuous_ones(r)) # 输出: [(0, 0), (3, 5), (7, 9)]
性能说明
该方法完全基于Numpy的底层向量化运算,避免了Python循环的开销,处理200万元素的数组时耗时可降至毫秒级,相比原方法有数量级的性能提升。
内容的提问来源于stack exchange,提问作者John
相关产品推荐
相关产品推荐

