Numpy张量中逐矩阵点积的优雅实现方法问询
优雅实现Numpy张量的逐矩阵点积操作
嘿,这个问题一点都不简单!其实Numpy里有好几种非常优雅的方式来实现你要的逐batch矩阵点积操作,我给你列几个最常用的:
方法1:使用np.matmul(推荐)
np.matmul是Numpy专门为矩阵乘法设计的函数,对于高维张量,它会自动将前面的维度视为batch维度,只对最后两个维度执行矩阵乘法,完美匹配你的需求:
import numpy as np M = np.array(range(12)).reshape(2,2,3) N = np.array(range(12)).reshape(2,3,2) # 逐batch矩阵点积 result = np.matmul(M, N) # 验证结果 print("result[0]等于np.dot(M[0], N[0]):", np.array_equal(result[0], np.dot(M[0], N[0]))) print("result[1]等于np.dot(M[1], N[1]):", np.array_equal(result[1], np.dot(M[1], N[1])))
输出结果的形状是(2,2,2),其中每个result[i]就是你要的np.dot(M[i], N[i])的结果。
方法2:使用@运算符(更简洁)
@是np.matmul的中缀语法糖,写法更简洁直观,功能完全一致:
result = M @ N
这种写法在代码里更清爽,适合日常使用。
方法3:使用np.einsum(灵活性拉满)
如果你需要更精细地控制维度运算,np.einsum是绝佳选择,它通过爱因斯坦求和约定来定义运算,可读性和灵活性都很强:
result = np.einsum('bij,bjk->bik', M, N)
这里的字符串'bij,bjk->bik'可以理解为:
b:batch维度(对应你的第一个维度,共2个batch)i,j:M的矩阵维度(2行3列)j,k:N的矩阵维度(3行2列)- 箭头后的
bik表示输出每个batch下的2行2列矩阵,正好是矩阵乘法的结果。
内容的提问来源于stack exchange,提问作者MH Ng
相关产品推荐
相关产品推荐

