如何在NumPy中实现元组风格的argsort排序?
在NumPy中实现元组式argsort的高效方法
我有一个由有序对组成的有序对数组:
import numpy as np foo = np.array([ [[1, 3], [2, 4]], [[4, 1], [2, 3]], [[3, 2], [3, 1]], [[2, 4], [3, 1]], [[1, 2], [3, 4]]])
我需要按照Python元组的常规排序规则,对每一组有序对执行argsort操作,期望得到的索引数组如下:
full_argsort(foo) == np.array( [[0, 1], [1, 0], [1, 0], [0, 1], [0, 1]])
(只要是能通过foo[indices]实现排序的索引数组都可以,目前不确定最优格式)
几个具体示例:
pair_argsort([[1,2],[2,4]]) == [0,1] pair_argsort([[2,1],[1,2]]) == [1,0] pair_argsort([[1,2],[1,1]]) == [1,0] # 注意:当第一个元素相同时,比较第二个元素
更通用的表述(转换为数组/元组/列表形式):
pair_of_pairs[pair_argsort(pair_of_pairs)] == sorted(map(tuple, pair_of_pairs))
我可以用np.apply_along_axis(pair_argsort,...)对整个foo执行argsort,但这种调用Python函数的方式效率很低,希望找到NumPy原生的实现方法。
我尝试过仅按元组的第一个元素做argsort,代码是np.argsort(foo.reshape(-1, 2, 2), axis=1)[:,:,0],得到的结果是:
array([[0, 1], [1, 0], [0, 1], [0, 1], [0, 1]])
但这个方法在处理第三组有序对[[3,2],[3,1]]时结果错误,因为它没有考虑第一个元素相同时的第二个元素比较。
请问如何在NumPy中实现符合元组排序规则的argsort?
内容的提问来源于stack exchange,提问作者Him
相关产品推荐
相关产品推荐

