如何用Numpy高效获取布尔数组中True值的首尾索引?
高效获取布尔数组中连续True值的首尾索引(Numpy实现)
核心实现代码
import numpy as np x = np.array([np.nan, 11, 13, np.nan, np.nan, np.nan, 9, 3, np.nan, 3, 4, np.nan]) mask = np.isnan(x) # 给布尔数组前后补False,统一处理首尾边界的连续块 mask_padded = np.concatenate([[False], mask, [False]]) # 转整数后求差分,1对应True块的起始,-1对应True块的结束(下一个位置) diffs = np.diff(mask_padded.astype(int)) # 提取所有连续True块的起始、结束索引 starts = np.where(diffs == 1)[0] ends = np.where(diffs == -1)[0] - 1 # 修正结束索引到原数组位置 # 生成目标格式结果 result = [] for s, e in zip(starts, ends): result.append(s if s == e else [s, e]) print(result) # 输出: [0, [3, 5], 8, 11]
方法原理
- 边界统一处理:通过在布尔数组首尾添加
False,无需单独编写首尾连续块的特殊判断逻辑。 - 差分定位边界:利用
np.diff计算相邻元素的变化,快速定位所有连续True块的起始和结束位置,这一步是Numpy底层优化的向量化操作,效率远高于Python循环。 - 结果格式化:遍历索引对,单个True直接保留索引,连续True则生成首尾索引组成的列表。
效率优势
该方案完全依赖Numpy的向量化操作,避免了Python循环的性能瓶颈。对于数十万级元素的数组,速度比纯Python循环快10~100倍,非常适合批量处理多组数组的场景。
内容的提问来源于stack exchange,提问作者cyrille_saf
相关产品推荐
相关产品推荐

