np.argsort按矩阵第3列排序,值超10时失效问题求助
问题
使用np.argsort对矩阵按第3列排序时出现异常:当矩阵中(3,3)位置元素大于10时,排序无法正常工作;该元素为9或更小值时,排序结果正确。
代码示例:
import numpy as np X = np.array([[5, 2, 3], [2.222, 5.5, 6], [3.3, 8, 10], [1.05, 0, 0]]) y = np.array([['T'], ['F'], ['T'], ['T']]) data = np.column_stack([X,y]) print(data) sorted_data = data[data[:, 2].argsort()] print(sorted_data)
原因及解决方法
核心原因
问题出在np.column_stack([X,y])这一步:因为y是字符串类型的数组,合并后整个data数组的 dtype 会被统一转为字符串类型。此时对第3列排序时,是按照字符串的字典序而非数值大小排序。
比如字符串"10"和"3"比较,字典序中"1"的ASCII码小于"3",所以"10"会被排在"3"前面,这就导致了数值上10>3但排序结果不符合预期的情况;而当元素是9时,字符串"9"的字典序大于"3"、"6"等,和数值排序结果一致,所以看起来正常。
解决方法
有两种常用的处理方式:
- 分开处理数值和标签:先对数值数组
X按第3列排序,再同步调整标签数组y的顺序:import numpy as np X = np.array([[5, 2, 3], [2.222, 5.5, 6], [3.3, 8, 10], [1.05, 0, 0]]) y = np.array([['T'], ['F'], ['T'], ['T']]) # 获取排序索引 sort_idx = X[:, 2].argsort() # 分别排序X和y sorted_X = X[sort_idx] sorted_y = y[sort_idx] # 如需合并,可使用结构化数组保持类型区分 sorted_data = np.rec.fromarrays([sorted_X[:,0], sorted_X[:,1], sorted_X[:,2], sorted_y[:,0]], names=('col1','col2','col3','label')) print(sorted_data) - 使用结构化数组:创建时明确各列的数据类型,避免类型统一转换:
import numpy as np # 创建结构化数组 data = np.array([(5, 2, 3, 'T'), (2.222, 5.5, 6, 'F'), (3.3, 8, 10, 'T'), (1.05, 0, 0, 'T')], dtype=[('col1', float), ('col2', float), ('col3', float), ('label', 'U1')]) # 按col3排序 sorted_data = np.sort(data, order='col3') print(sorted_data)
内容的提问来源于stack exchange,提问作者ebhh
相关产品推荐
相关产品推荐

