同维度Numpy 3D矩阵乘法方法及np.einsum使用咨询
3D矩阵乘法与np.einsum用法解析
一、3D矩阵的"类转置乘法"实现
对于两个形状均为(M, N, D)的3D矩阵,你可以把它们看作D个独立的2D矩阵集合(每个位置d对应一个M×N的2D矩阵)。类似2D中通过转置实现同形状矩阵的乘法,3D场景下可以批量处理每个D维度的切片:
场景1:每个切片执行 A[:, :, d] @ B[:, :, d].T(输出形状(M, M, D))
方法1:轴调整 + np.matmul
通过调整轴的顺序,把D维度放到批量维度的位置,再进行矩阵乘法,最后调整回原维度顺序:
import numpy as np # 生成示例矩阵 M, N, D = 2, 3, 4 A = np.random.rand(M, N, D) B = np.random.rand(M, N, D) # 调整轴并计算 result = np.matmul(A.transpose(2, 0, 1), B.transpose(2, 1, 0)).transpose(1, 2, 0)
A.transpose(2,0,1):把A从(M,N,D)转为(D,M,N),将D作为批量维度B.transpose(2,1,0):把B从(M,N,D)转为(D,N,M),等价于每个2D切片转置- 批量矩阵乘法后得到
(D,M,M),再转置为(M,M,D)
方法2:直接用np.einsum(更直观)
result = np.einsum('mnd, nmd -> mmd', A, B)
场景2:每个切片执行 A[:, :, d].T @ B[:, :, d](输出形状(N, N, D))
方法1:轴调整 + np.matmul
result = np.matmul(A.transpose(2, 1, 0), B.transpose(2, 0, 1)).transpose(1, 2, 0)
方法2:np.einsum实现
result = np.einsum('nmd, mnd -> nnd', A.swapaxes(0,1), B)
二、np.einsum的工作原理与用法
核心原理
np.einsum基于爱因斯坦求和约定,通过字符串直接定义维度的运算逻辑:
- 字符串中用逗号分隔多个输入的维度标识(每个字母代表一个维度)
- 重复出现在多个输入中的字母,会自动对该维度进行求和(相当于矩阵乘法中的"收缩"维度)
- 箭头
->后面的字母表示输出保留的维度及顺序
举例子理解
2D矩阵乘法:
A(M×N) @ B(N×P) = C(M×P),对应einsum写法:np.einsum('mn, np -> mp', A, B)这里
n是重复维度,会被求和,最终保留m和p维度。回到你的3D场景:
np.einsum('mnd, nmd -> mmd', A, B)- 第一个输入
A的维度:m(行)、n(列)、d(批量) - 第二个输入
B的维度:n(行,对应A的列)、m(列,对应B切片的转置)、d(批量) - 重复维度
n会被求和,最终保留m(A的行)、m(B的列)、d(批量),输出形状(M,M,D)
- 第一个输入
优势
相比轴调整+matmul,np.einsum不需要手动调整维度顺序,直接通过维度标识就能清晰表达运算逻辑,尤其适合复杂的多维张量运算。
内容的提问来源于stack exchange,提问作者user18980990
相关产品推荐
相关产品推荐

