如何高效实现Numpy中np.diagonal(np.dot(A,B),axis1=1,axis2=2)操作?
高效计算Numpy数组指定对角线元素的方案
你提到的np.diagonal(np.dot(A, B), axis1=1, axis2=2)确实存在大量冗余计算——np.dot(A, B)会生成形状为(n, m, m)的完整矩阵,但我们只需要每个子矩阵的主对角线元素,完全没必要计算所有非对角线项。下面是两种更高效的替代方法:
方法1:使用np.einsum
einsum可以直接指定维度间的求和关系,精准计算我们需要的对角线元素,避免冗余:
import numpy as np result = np.einsum('nmk,kj->nj', A, B)
这里的下标含义:
nmk对应数组A的(n, m, k)维度kj对应数组B的(k, m)维度->nj表示对k维度求和,最终得到(n, m)的结果,和原方法输出完全一致。
方法2:广播相乘后求和
通过广播机制让A和B的对应元素相乘,再沿k维度求和,同样能得到目标结果:
result = (A * B.T[None, ...]).sum(axis=2)
步骤拆解:
- 将B转置为
(m, k),再添加一个维度变成(1, m, k),和A的(n, m, k)实现广播匹配 - 对应元素相乘得到
(n, m, k)的数组 - 沿
k维度求和,得到(n, m)的对角线结果
验证与性能对比
用随机数组验证结果一致性:
n, m, k = 2, 3, 4 A = np.random.rand(n, m, k) B = np.random.rand(k, m) original = np.diagonal(np.dot(A, B), axis1=1, axis2=2) einsum_res = np.einsum('nmk,kj->nj', A, B) broadcast_res = (A * B.T[None, ...]).sum(axis=2) print(np.allclose(original, einsum_res)) # 输出True print(np.allclose(original, broadcast_res)) # 输出True
性能上,两种新方法的时间复杂度都是O(n*m*k),远低于原方法的O(n*m²*k),当m较大时,效率提升会非常明显。
内容的提问来源于stack exchange,提问作者Yandle
相关产品推荐
相关产品推荐

