旧版NumPy 1.10中高效获取分组指定列最小值行的方法
高效处理旧版NumPy下的分组取最小值行问题
嘿,太懂这种旧版库卡脖子的感觉了!20万行的数据用循环遍历唯一组合确实会慢到让人崩溃,咱们换个排序+分组标记的思路,能把速度提上去不止一个档次,而且完全兼容NumPy 1.10。
核心思路
既然没办法直接用np.unique(axis=0),那咱们先把数据按分组列(1、2、4)排序,再按目标列(0)升序排序——这样同一分组的行都会挨在一起,而且每个分组里列0最小的行肯定排在该组的最前面。之后只要找出每个分组的第一行,就是咱们要保留的结果。
具体实现代码
假设你的数据存在arr里,shape是(200000, 6),步骤如下:
- 按规则排序
先按分组列(1、2、4)排序,再按列0升序排序,确保同一分组内最小的列0行在最前面:
# lexsort的参数是(次要排序键, 主要排序键...),这里先按列0,再按1、2、4 sorted_indices = np.lexsort((arr[:, 0], arr[:, 1], arr[:, 2], arr[:, 4])) sorted_arr = arr[sorted_indices]
- 生成分组掩码
接下来要标记每个分组的第一行。因为NumPy 1.10不支持多维unique,咱们可以把分组列转成结构化数组(或者用视图转成一维复合键),这样就能直接比较相邻行是否属于同一分组:
# 提取分组列 group_cols = sorted_arr[:, [1, 2, 4]] # 转成结构化数组,dtype要和原数据匹配,比如原数据是float64就用这个 structured_groups = np.rec.fromarrays( [group_cols[:,0], group_cols[:,1], group_cols[:,2]], dtype=[('col1', group_cols.dtype), ('col2', group_cols.dtype), ('col4', group_cols.dtype)] ) # 生成掩码:第一行必选,后续行如果和前一行分组不同则选 mask = np.concatenate([[True], structured_groups[1:] != structured_groups[:-1]]) # 筛选结果 result = sorted_arr[mask]
如果你的分组列都是同一种数值类型(比如都是float或int),还可以用更高效的视图方法:
# 将三列转成一维复合键视图,dtype格式为(原类型, 列数) group_view = sorted_arr[:, [1,2,4]].view(dtype=(sorted_arr.dtype, 3)).ravel() # 生成掩码 mask = np.concatenate([[True], group_view[1:] != group_view[:-1]]) result = sorted_arr[mask]
为什么这个方法快?
原来的循环方法时间复杂度是O(n*k)(k是唯一分组数),而这个方法的时间复杂度主要来自排序的O(n log n),后续的掩码生成是O(n)——对于20万行的数据,排序在NumPy里是高度优化的C实现,速度会比循环快几个数量级。
注意事项
- 如果分组列里有
NaN,要特别注意:因为NaN != NaN,所以每个含NaN的行会被当成单独的分组。如果需要把NaN视为同一组,得先把NaN替换成一个特定值(比如某个极端值)再处理。 - 确保结构化数组的dtype和原数据一致,否则会出现类型转换错误。
内容的提问来源于stack exchange,提问作者cataclysmic
相关产品推荐
相关产品推荐

