如何用NumPy实现三维数组的点积运算?
如何用NumPy实现三维数组的点积运算?
嘿,我来帮你搞定这个三维点积的问题!先看看你遇到的状况:你想把二维里的点积逻辑扩展到三维数组上,但用np.dot时因为维度不匹配报错了,对吧?
首先咱们理清楚错误根源:你的W形状是(2,2,3),x是(2,1,3),np.dot对高维数组的运算规则是用第一个数组的最后一维和第二个数组的倒数第二维做乘法,这里W的最后一维是3,x的倒数第二维是1,两者不相等,自然就报错啦。
再回忆你那行得通的二维例子:W是(2,3),x是(3,),np.dot(W,x)本质是让W的每一行和x做向量点积,得到(2,)的结果。放到三维场景里,你应该是想让每个batch(也就是第一个维度的2个样本)里,各自的(2,3)的W和(1,3)的x做和二维一样的点积,最终得到每个batch对应(2,)的结果,对吧?
那接下来给你几个靠谱的解决方案:
方案1:用np.einsum(最直观,精准控制维度运算)
einsum可以通过字符串指定维度之间的运算关系,非常适合这种需要明确对应维度的场景:
import numpy as np x = np.array([[[-1, 2, -4]], [[-1, 2, -4]]]) W = np.array([[[2, -4, 3], [-3, -4, 3]], [[2, -4, 3], [-3, -4, 3]]]) # 解释:b代表batch维度(2个),i是W的行维度(2行),j是特征维度(3个特征) # 实现每个batch里,W的第i行和x的行在j维度做点积 y = np.einsum('bij,bkj->bik', W, x) # 去掉多余的单维度,得到更简洁的结果 y_clean = np.squeeze(y) print(y_clean)
运行后会输出:
[[-22 -17] [-22 -17]]
正好对应每个batch里和二维例子一致的结果~
方案2:用np.matmul(或@运算符,代码更简洁)
matmul会自动处理batch维度,只要前面的batch维度匹配,我们只需要调整x的维度让它和W的最后一维对齐就行:
import numpy as np x = np.array([[[-1, 2, -4]], [[-1, 2, -4]]]) W = np.array([[[2, -4, 3], [-3, -4, 3]], [[2, -4, 3], [-3, -4, 3]]]) # 把x转成(2,3,1),让W的最后一维(3)和x的倒数第二维(3)匹配 x_reshaped = x.transpose(0, 2, 1) # @运算符等价于np.matmul y = W @ x_reshaped y_clean = np.squeeze(y) print(y_clean)
这个方案也能得到和上面一样的结果,代码更简洁直观~
方案3:手动遍历batch(适合新手理解,效率稍低)
如果你想更直观地看到每个batch的运算过程,可以用列表推导式逐个处理每个batch:
import numpy as np x = np.array([[[-1, 2, -4]], [[-1, 2, -4]]]) W = np.array([[[2, -4, 3], [-3, -4, 3]], [[2, -4, 3], [-3, -4, 3]]]) # 把每个batch的x转成(3,)的向量,和二维例子里的x形状一致 y = np.array([np.dot(w_batch, x_batch.reshape(3,)) for w_batch, x_batch in zip(W, x)]) print(y)
这个方法虽然易懂,但数据量较大时效率不如前两种向量化方法,更推荐前两种哦。
备注:内容来源于stack exchange,提问作者Sun Bear
相关产品推荐
相关产品推荐

