PyTorch中最后维度不同的两个张量如何实现乘法运算?
PyTorch中不同形状张量的@运算符(批量矩阵乘法)运算逻辑
首先明确:PyTorch中的@运算符实现的是支持广播的批量矩阵乘法,等价于torch.matmul,不是逐元素乘法(逐元素用*),也不是普通的二维矩阵乘法,它的核心规则是:
- 对于张量
x(形状(..., m, n))和y(形状(..., n, p)),x@y的结果形状为(..., m, p) - 这里的
...表示任意数量的前置“批量维度”,这些维度会自动按广播规则对齐
你的例子拆解
来看你给出的张量:
import torch a = torch.arange(0,9).view(3,3) # 形状: (3,3) → 等价于隐含批量维度为1的(1,3,3) b = torch.arange(0,30).view(2,3,5) # 形状: (2,3,5)
步骤1:广播对齐批量维度
a的原始形状是(3,3),没有前置批量维度;b的前置批量维度是(2,)。根据广播规则,a会被自动扩展为(2,3,3)(相当于在第0维度复制一次),和b的批量维度对齐。
步骤2:逐批量执行矩阵乘法
广播完成后,a的每个批量样本是(3,3)的矩阵,b的每个批量样本是(3,5)的矩阵。对每个批量位置,执行标准的二维矩阵乘法:(3,3) @ (3,5) → 得到(3,5)的结果。
步骤3:保留批量维度输出
所有批量的结果组合起来,最终输出形状就是(2,3,5),这和PyTorch的实际输出一致。
手动验证一个元素
比如取输出张量的[0,0,0]位置:
a的第一行是[0,1,2]b的第一个批量的第一列是[0,5,10]- 点积计算:
0*0 + 1*5 + 2*10 = 25
你可以运行代码验证:
result = a @ b print(result[0,0,0]) # 输出25,和手动计算一致
常见误解纠正
- 不要把
@和逐元素乘法*混淆:逐元素乘法要求所有维度完全匹配(或可广播到完全匹配),而@只要求最后两个维度满足矩阵乘法的维度条件(第一个张量的最后一维=第二个张量的倒数第二维)。 - 前置维度是独立批量,不是要转置后相乘:每个批量内的矩阵是独立计算的,不会跨批量操作,所以最终保留批量维度。
内容的提问来源于stack exchange,提问作者Hossein
相关产品推荐
相关产品推荐

