如何使用Numpy高效实现按索引矩阵对目标矩阵逐行索引
高效实现Numpy矩阵按行自定义索引重排
要实现根据每行的列索引矩阵ind重排矩阵A的每行(满足B[i,j] = A[i, ind[i,j]]),避免Python循环带来的性能损耗,最优方案是使用Numpy的高级索引机制,完全依托Numpy内置的C级操作提升速度。
核心实现代码
import numpy as np # 示例输入矩阵A A = np.array([[88, 44, 77, 33, 77], [33, 55, 66, 88, 0], [88, 0, 0, 55, 88], [ 0, 22, 44, 88, 33], [33, 33, 77, 66, 66]]) # 每行的列索引矩阵(示例用argsort生成) ind = np.argsort(A) # 构造行索引数组:将一维行索引转为(n_rows, 1)的二维数组,和ind的维度广播匹配 row_indices = np.arange(A.shape[0])[:, None] # 直接通过高级索引得到结果矩阵B B = A[row_indices, ind] print(B)
输出结果:
array([[33, 44, 77, 77, 88], [ 0, 33, 55, 66, 88], [ 0, 0, 55, 88, 88], [ 0, 22, 33, 44, 88], [33, 33, 66, 66, 77]])
适配ind列数少于A的场景
如果ind的列数少于A,该方法依然直接适用,无需修改逻辑:
# 取每行前3个排序索引 ind_short = ind[:, :3] B_short = A[row_indices, ind_short] print(B_short)
输出结果:
array([[33, 44, 77], [ 0, 33, 55], [ 0, 0, 55], [ 0, 22, 33], [33, 33, 66]])
性能优势说明
原方案的列表推导式本质是Python层面的循环,而高级索引的操作完全在Numpy内部以C语言执行,没有Python解释器的开销。当矩阵规模较大时(比如1000×1000及以上),这种实现的速度会比循环方案快几十倍甚至上百倍,完全满足对性能要求高的业务场景。
内容的提问来源于stack exchange,提问作者Tyler
相关产品推荐
相关产品推荐

