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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 01:29:52