高维numpy数组执行指定维度点积的最优实现方案咨询
NumPy高维指定轴点积最优实现
你当前的循环方案为Python层级的遍历运算,数组规模较大时性能较低,可使用以下完全矢量化的原生NumPy方案实现需求,运算效率远高于循环实现:
方案1:使用np.einsum(最推荐,无需额外调整维度)
einsum支持通过下标标记直接定义维度对应关系和收缩规则,写法最直观:
import numpy as np # 测试数组 a = np.random.rand(20,10,6,5,4) b = np.random.rand(10,6,5,4,3) # 直接按规则计算,输出形状完全匹配预期 res = np.einsum('ijklm,jklmn->inklm', a, b) print(res.shape) # 输出 (20, 3, 6, 5, 4)
方案2:使用np.tensordot + 轴调整
通过axes参数指定需要收缩求和的维度,运算后调整轴顺序即可得到目标结果:
# 收缩a的第1维与b的第0维,再将最后一维的3移动到第1位 res = np.moveaxis(np.tensordot(a, b, axes=([1], [0])), -1, 1) print(res.shape) # 输出 (20, 3, 6, 5, 4)
方案3:使用广播矩阵乘法(NumPy 1.16+支持)
将公共批次维度调整到最前,即可直接用@运算符批量执行矩阵乘法:
a_trans = a.transpose(2,3,4,0,1) # 形状变为 (6,5,4,20,10) b_trans = b.transpose(1,2,3,0,4) # 形状变为 (6,5,4,10,3) res = (a_trans @ b_trans).transpose(3,4,0,1,2) print(res.shape) # 输出 (20, 3, 6, 5, 4)
你可以通过np.allclose验证以上三种方案的输出结果与你原有循环实现的结果完全一致。
内容的提问来源于stack exchange,提问作者Mert Onur
相关产品推荐
相关产品推荐

