如何优化Numpy中Argmax向量点积的索引提取以避免冗余?
优化向量点积argmax后的索引提取代码
问题背景
你有两组向量集合:
gamma是形状为(m,n,a)的数组,对应结构[[v11,...,v1n],...,[vm1,...,vmn]],维度标识为(i,j,s)b是形状为(k,a)的数组,对应结构[b1,...,bk],维度标识为(k,s)
需求是对所有 i,x 计算 argmax_j dot(vij, bx),当前实现通过tensordot得到(i,j,k)的点积数组,取argmax得到(i,k)的索引后,用np.take提取时生成了冗余的(i,i,k,s)数组,最后靠取对角线得到结果,希望避免这一冗余步骤。
优化方案
可以利用numpy的高级索引直接定位所需元素,完全跳过生成冗余中间数组的步骤。核心思路是为argmaxes补充对应i维度的索引,直接从gamma中提取每个i对应的最优j向量。
优化后代码
import numpy as np # 确保输入为numpy数组(若原始数据不是的话) gamma = np.array(gamma) # shape (m, n, a) b = np.array(b) # shape (k, a) # 计算点积并得到argmax索引,结果shape为 (m, k) argmaxes = np.tensordot(gamma, b, axes=[[2], [1]]).argmax(axis=1) # 生成i维度的索引数组,shape为 (m, 1),用于和argmaxes广播匹配 i_indices = np.arange(gamma.shape[0])[:, np.newaxis] # 直接通过高级索引提取目标结果,最终shape为 (m, k, a) result = gamma[i_indices, argmaxes]
原理说明
i_indices = np.arange(m)[:, np.newaxis]生成形状为(m,1)的索引数组,和argmaxes(形状(m,k))广播后,会为每个(i,k)位置匹配对应的i索引。gamma[i_indices, argmaxes]利用numpy的高级索引特性,直接定位每个i下对应argmaxes[i,k]的j向量,一步得到目标形状的结果,完全不会生成原方法中(m,m,k,a)的冗余大数组,内存占用和计算效率都显著提升。
与原方法对比
原方法中np.take(gamma, argmaxes, axis=1)会把每个i对应的argmaxes[i,:]索引应用到所有i维度上,导致生成包含大量重复数据的(m,m,k,a)数组,后续取对角线的操作完全是在清理冗余数据。而高级索引直接针对每个i取对应的j,从根源上避免了冗余计算。
内容的提问来源于stack exchange,提问作者Daniel Robert-Nicoud
相关产品推荐
相关产品推荐

