如何用NumPy高效计算遇0重置的连续1累积和?
高效计算NumPy数组中连续1的累积和(遇0重置)
针对大规模0-1二维NumPy数组,循环实现连续1的累积计数(遇0重置)效率低下,这里提供完全矢量化的高效方案,无需显式循环:
实现代码
import numpy as np arr_in = np.array([[1,1,1,1,1,1], [0,0,0,0,0,0], [1,0,1,0,1,1], [0,1,1,1,0,0]]) # 1. 计算每行的累积和 cumulative_sum = arr_in.cumsum(axis=1) # 2. 标记所有0的位置,仅保留这些位置的累积和,其余设为0 reset_markers = np.where(arr_in == 0, cumulative_sum, 0) # 3. 对每行的重置标记取累积最大值,得到每个位置最近一次0的累积和 last_reset = np.maximum.accumulate(reset_markers, axis=1) # 4. 用整体累积和减去最近一次重置的累积和,得到目标结果 arr_out = cumulative_sum - last_reset print(arr_out)
输出结果
[[1 2 3 4 5 6] [0 0 0 0 0 0] [1 0 1 0 1 2] [0 1 2 3 0 0]]
原理说明
- 利用
cumsum(axis=1)快速计算每行的整体累积和,这一步是矢量化操作,效率极高。 - 通过
np.where标记所有0的位置,记录这些位置的累积和值,作为重置点标记。 np.maximum.accumulate会沿每行向前传播最近的重置点累积和,确保每个位置都能获取到上一次遇到0时的累积值。- 最后用整体累积和减去最近重置点的累积值,自动实现连续1的计数重置,完全符合需求。
这种方案全程基于NumPy的矢量化运算,避免了Python循环的开销,能够轻松处理数万行/列的大规模数组。
内容的提问来源于stack exchange,提问作者tibibou
相关产品推荐
相关产品推荐

