如何使用tensordot实现带额外维度的批量矩阵乘法
带批量维度的二维矩阵乘法实现方案
你需要实现的是沿第三维度逐通道独立计算二维矩阵乘积,即对每个通道下标k,单独计算a[:,:,k] @ b[:,:,k],再将所有通道的结果拼接为(2,2,3)的输出。
错误原因说明
- 原
tensordot(a,b,1)默认收缩a的最后一维、b的第一维,当a最后一维为3、b第一维为2时,维度尺寸不匹配触发报错 - 配置
axes=((0,1),(0,1))会完全收缩a、b的前两维,仅保留双方的第三维,输出形状为(3,3),不符合预期
实现方案
方案1:使用np.einsum(最简洁)
直接通过爱因斯坦求和约定指定维度运算规则,无需调整轴顺序:
import numpy as np # 生成测试输入 a = np.random.rand(2,2,3) b = np.random.rand(2,2,3) # 按规则计算乘积 res_einsum = np.einsum('ijk,jlk->ilk', a, b) print(res_einsum.shape) # 输出 (2, 2, 3)
方案2:调整轴顺序后使用@运算符(可读性更高)
numpy的@运算符原生支持批量维度运算,只需将批量通道维度调整到最前:
# 将第三维移到最前,转为(3,2,2)的批量矩阵格式 a_batch = np.moveaxis(a, -1, 0) b_batch = np.moveaxis(b, -1, 0) # 批量矩阵乘法,结果形状为(3,2,2) res_batch = a_batch @ b_batch # 将批量维度移回最后,得到(2,2,3)的输出 res_at = np.moveaxis(res_batch, 0, -1) print(res_at.shape) # 输出 (2, 2, 3)
结果验证
两种推荐方案的输出完全一致:
print(np.allclose(res_einsum, res_at)) # 输出 True
内容的提问来源于stack exchange,提问作者Brandon Dube
相关产品推荐
相关产品推荐

