You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.06 08:52:50