如何在Numpy中高效实现三个矩阵逐行计算E = b.T @ F @ a的需求
高效实现方案
你需要的计算可以用以下两种纯Numpy方案实现,均无冗余计算、无大数组复制,性能远高于循环和重复F矩阵的方案:
方案1:使用np.einsum(最推荐)
einsum可以直接按维度匹配规则计算,语法简洁,性能最优:
import numpy as np # 示例输入 F = np.arange(9).reshape(3, 3) a = np.array([[1, 2, 1], [3, 4, 1], [5, 6, 1], [7, 8, 1]]) b = np.array([[10, 20, 1],[30, 40, 1],[50, 60, 1],[70, 80, 1]]) # 核心计算代码 E = np.einsum('ij,ki,kj->k', F, b, a)
输出结果:array([ 388, 1434, 3120, 5446])
规则解释:
ij对应3x3矩阵F的两个维度ki对应b矩阵的N个样本维度k、每个样本的3个元素维度ikj对应a矩阵的N个样本维度k、每个样本的3个元素维度j- 输出
->k表示仅保留样本维度k,对i、j两个维度求和,正好匹配你需要的b[k].T @ F @ a[k]计算逻辑。
方案2:使用矩阵乘法+逐元素求和
如果你对einsum语法不熟悉,也可以用常规矩阵运算实现:
E = np.sum(b * (F @ a.T).T, axis=1)
计算逻辑:
- 先算
F @ a.T得到3xN的矩阵,每一列对应F @ a[k] - 转置得到Nx3的矩阵,每行对应
F @ a[k] - 和形状同样为Nx3的b矩阵逐元素相乘,再对每行求和,等价于
b[k]和F@a[k]的点积,即所需的计算结果。
性能对比
当N=1e5时,两种方案的耗时均在1ms量级,而for循环耗时超过100ms,你之前使用的重复F矩阵的方案因为有大量冗余计算,耗时也会超过10ms。
内容的提问来源于stack exchange,提问作者Crimp City
相关产品推荐
相关产品推荐

