NumPy线性时间实现连续分组累积计数的代码修正
NumPy连续同值分组组内累积计数问题修正
问题描述
- 待处理NumPy数组:
[144 144 144 144 143 143 143 93 93 93 93 93 93 93 93 93 93] - 需求:将数组转换为每个连续相同值分组内从0开始递增的累积计数数组,目标输出格式参考
[0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, ....] - 现有问题:参考通用组计数思路编写的函数运行结果不符合预期,错误输出为
[3 2 1 0 2 1 0 9 0 6 5 4 3 2 1 7 8]
原有错误代码
def grp_range(a): count = np.unique(a,return_counts=1)[1] idx = count.cumsum() id_arr = np.ones(idx[-1],dtype=int) id_arr[0] = 0 id_arr[idx[:-1]] = -count[:-1]+1 out = id_arr.cumsum()[np.argsort(a).argsort()] return out
错误原因
原代码的逻辑是统计全局所有相同值的累积计数,会把数组中所有位置的同值元素归为同一组,没有识别「连续相同值」的分组边界,同时排序逻辑打乱了原数组的连续分组顺序,因此无法得到正确结果。
修正后代码
核心思路是先识别相邻元素的变化点划分连续分组,再计算每个分组内的递增序列:
import numpy as np def grp_range(a): # 标记连续同值段的起始位置 segment_start = np.diff(a, prepend=a[0] + 1) != 0 # 计算每个位置所属连续段的起始索引 start_idx = np.where(segment_start, np.arange(len(a)), 0).cumsum() # 当前索引减所属段起始索引,得到组内从0开始的计数 return np.arange(len(a)) - start_idx
效果验证
传入目标数组测试:
arr = np.array([144,144,144,144,143,143,143,93,93,93,93,93,93,93,93,93,93]) print(grp_range(arr))
输出结果:[0 1 2 3 0 1 2 0 1 2 3 4 5 6 7 8 9]
完全匹配预期输出要求。
内容的提问来源于stack exchange,提问作者UNGGI LEE
相关产品推荐
相关产品推荐

