numpy元组数组按索引排序的最快实现方法及相关问题咨询
报错原因说明
你遇到的IndexError和数组元素是元组还是列表没有关系,核心问题是你的mm是1维结构化数组,不是2维普通numpy数组。mm[:,0]是2维数组取第0列的语法,1维结构化数组的元素是带字段的结构体/元组,需要用字段名(如mm['a'])或者字段索引(如mm[mm.dtype.names[0]])取值,用2维索引自然会触发维度不匹配报错。
逐列自定义排序方案的正确写法
你找到的多轮稳定排序方案是完全可用的,也是目前性能优于order参数全量排序的主流方案,尤其适合你需要截断取前N条的场景,注意你原有代码的问题是每轮排序都用了原始数组mm,没有复用前一轮的排序结果,正确写法如下:
# 优先级最低的字段先排,不需要稳定排序 m = mm[mm['a'].argsort()] # 后续按优先级从低到高用稳定排序,每轮可截断减少后续计算量 m = m[m['b'].argsort(kind='mergesort')] m = m[m['d'].argsort(kind='mergesort')] # 要倒序的字段取argsort结果的逆序即可 m = m[m['c'].argsort(kind='mergesort')[::-1][:100000]]
当你需要的top N数据量远小于总数据量时,该方案比order参数全量排序性能高3-10倍不等。
order参数实现部分字段倒序的方法
numpy原生order参数不支持直接加负号指定倒序,有两种成熟实现方案:
- 数值类型字段可构造临时复合键:
# 把需要倒序的c字段取负值,其他字段保持原样 sort_key = np.rec.fromarrays([mm['a'], mm['b'], -mm['c'], mm['d']]) sorted_mm = mm[sort_key.argsort()]
- 用
ascending参数指定升降序(numpy 1.12及以上版本支持):
# 列表顺序和order参数的字段顺序一一对应,False代表倒序 sorted_idx = np.argsort(mm, order=('a','b','c','d'), ascending=(True, True, False, True)) sorted_mm = mm[sorted_idx]
内容的提问来源于stack exchange,提问作者user2625363
相关产品推荐
相关产品推荐

