You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何高效实现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)

步骤拆解:

  1. 将B转置为(m, k),再添加一个维度变成(1, m, k),和A的(n, m, k)实现广播匹配
  2. 对应元素相乘得到(n, m, k)的数组
  3. 沿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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.18 19:20:47