PyTorch 入门最难的不是安装,而是张量运算规则。很多人跑通了一个图像分类 Demo,然后自己写数据预处理时,一遇到逐元素计算、矩阵乘法、广播机制就开始报错:尺寸对不上、输出形状莫名多了一维、矩阵乘写成逐元素乘得到一堆标量。本课就把这三个规则拆开讲清楚,用最直白的方式说明它们分别适用什么场景、底层怎么对齐、报错时如何排查。
如果你是刚搭好 PyTorch 环境、正在从“创建一个 tensor”走向“能用张量做实际计算”的初学者,这一课正好处在承上启下的位置。后面学线性层、注意力机制、反向传播、卷积,本质上都逃不开这三类运算。先把它们吃透,很多看起来复杂的网络结构,拆到最后都是逐元素计算、矩阵乘法和广播的组合。
好,先不要急着堆代码,按顺序走。我们先确认环境,再讲规则,最后用一个综合例子把三个规则串起来。
1. 张量运算之前,先把环境问题确认清楚
1.1 为什么要先检查环境
很多人在本课卡住,不是不理解运算法则,而是环境没弄对。之前在网上能看到不少关于 pytorch 安装、pytorch 环境搭建、pytorch 下载很慢、pytorch 适配的问题。说实话,这些安装问题确实是第一道门槛,但不要让它变成今天的主线。我的建议是:如果你还在装环境,不要一上来就追 GPU 版。先装 CPU 版,足够跑完本课所有示例;等你后面真正要训练模型了,再按官方命令配置带 CUDA 的版本,这样能省很多折腾时间。
判断环境是不是正常,不需要看太多东西。只要在终端或 Python 解释器里执行下面几行:
import torch print(torch.__version__) print(torch.cuda.is_available())第一行能输出版本号,说明 PyTorch 核心库安装成功。第二行输出 True,说明当前环境能调用 CUDA;输出 False,说明当前只能使用 CPU,或者 GPU 驱动没匹配好。本课讲的是张量运算规则,CPU 和 GPU 在运算逻辑上完全一致,差异只在速度。所以即使输出 False,也不影响继续学习。
1.2 怎么处理下载慢和版本混乱
官方源如果速度不行,可以使用国内镜像源安装。配置时最需要注意的是版本一致性:你装的 PyTorch 版本、Python 版本、CUDA 版本三者要匹配。如果某个版本在官网找不到对应组合,就不要硬装,直接换一个稳定组合。
我见过不少同学用 conda 安装装了一半,又换成 pip 安装,结果环境里出现两个 PyTorch 版本。这种问题最隐蔽:表面上看安装成功,但 import 时可能加载了错误版本,或者在后续运行矩阵乘法时出现奇怪的报错。如果遇到这类问题,先把这个虚拟环境删掉重建,比反复尝试修复更容易解决问题。
1.3 快速跑通一个“张量运算冒烟测试”
环境确认之后,建议先跑一个冒烟测试,不做复杂操作,只验证张量创建和基本运算是否正常:
import torch a = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) b = torch.tensor([[5.0, 6.0], [7.0, 8.0]]) print(a + b) print(a * b)如果能正常输出两个 2x2 张量,说明环境基本可用。这里的a * b是逐元素乘法,它和矩阵乘法完全不同,下一章专门展开。
注意:前几课用 CPU 版完全没问题。真正需要 GPU 的时候,是后续训练较大的模型、处理较大 batch 时。到那时候再翻安装文档,会比一开始就追求 GPU 顺利得多。
2. 逐元素计算:同位置变量之间的一对一运算
2.1 什么是逐元素计算
逐元素计算是指两个张量在对应位置上的元素分别进行计算。简单说,就是“一对一”:a 的第 0 个元素和 b 的第 0 个元素算一次,a 的第 1 个元素和 b 的第 1 个元素算一次,其他位置同理。
为什么它是基础?因为深度学习里最频繁的操作,比如 ReLU 激活、加偏置、归一化、损失函数中的平方差,本质上都是逐元素计算。你不需要一开始就理解复杂的神经网络,只要能把两个同形状张量的逐元素加减乘除弄明白,后面很多公式就很自然了。
2.2 常用逐元素操作
用代码看最直观:
import torch a = torch.tensor([1.0, 2.0, 3.0]) b = torch.tensor([4.0, 5.0, 6.0]) print(a + b) # tensor([5., 7., 9.]) print(a - b) # tensor([-3., -3., -3.]) print(a * b) # tensor([ 4., 10., 18.]) print(a / b) # tensor([0.2500, 0.4000, 0.5000]) print(torch.pow(a, 2)) # tensor([1., 4., 9.]) print(a > 1) # tensor([False, True, True]) print(torch.exp(a)) # tensor([ 2.7183, 7.3891, 20.0855])其中最容易被误会的是*。在数学里我们经常把*当作乘法,但在 PyTorch 中,*默认是逐元素乘法,不是矩阵乘法。如果两个张量形状相同,*会对每一位对应相乘。如果你想做矩阵乘法,要用@或torch.matmul。
除四则运算外,比较运算也属于逐元素运算。a > 1返回一个形状相同的布尔张量,每个位置表示该位置是否满足条件。这在写掩码、做条件筛选的时候非常有用。
2.3 函数式操作和就地操作的区别
PyTorch 里同一类逐元素操作往往有两种写法。比如加法,可以写成a + b,也可以写成torch.add(a, b),两者返回的都是新张量。还有一类以_结尾的方法,比如a.add_(2),它会直接修改原张量。
就地操作有时能省内存,但也容易带来副作用。
c = torch.tensor([1.0, 2.0, 3.0]) c.add_(2) print(c) # tensor([3., 4., 5.])如果原张量还参与后续计算,或在自动求导中被需要,就地操作可能改变计算图,导致梯度计算异常。建议初学阶段尽量少用就地操作,多使用返回新张量的写法,逻辑更清晰。
2.4 逐元素计算容易踩的三个坑
第一,形状不一致直接相加会报错。比如[3]和[4]相加,PyTorch 不会自动把元素对齐,因为没有逐元素对应的关系。注意,它也不是完全不能处理不同形状,那要满足广播机制,下一章讲。
第二,把逐元素乘法和矩阵乘法混用。同一个*在 PyTorch 里不等于数学上的矩阵乘法,这是新手最容易看错的地方。
第三,不注意张量的 dtype。整型张量和浮点型张量做除法时,结果可能不是你想要的。比如torch.tensor([1, 2]) / torch.tensor([2, 2])在 Python 新版本中可能得到浮点,但在某些类型组合下会报错或截断。所以运算前尽量确认输入是浮点型,例如torch.tensor([1.0, 2.0])或调用.float()。
3. 矩阵乘法:内维对齐是唯一规则
3.1 矩阵乘法和逐元素乘法的本质区别
矩阵乘法不是一对一相乘,而是“行列相乘再相加”。
以二维矩阵为例,一个形状为(n, k)的矩阵乘以一个形状为(k, m)的矩阵,结果是(n, m)。中间的k是内维,必须相等。很多初学者记不住,我建议你把它写成:
(n, k) @ (k, m) -> (n, m)所以,如果遇到(3, 4)和(4, 5),结果就是(3, 5)。如果遇到(3, 4)和(3, 5),矩阵乘法会直接报错,因为第一个矩阵的列数 4 不等于第二个矩阵的行数 3。这时候你可能需要把其中一个矩阵转置。
3.2 最常用的矩阵乘法接口
PyTorch 里常用的矩阵乘法接口有这些:
| 接口 | 适用维度 | 说明 |
|---|---|---|
x @ y | 通用 | 语法糖,推荐日常使用 |
torch.matmul(x, y) | 通用 | 功能等同于@,支持广播 |
torch.mm(x, y) | 仅二维 | 比matmul更严格,少用 |
torch.bmm(x, y) | 仅三维批量 | 要求 batch 维度一致 |
先看最普通的二维矩阵乘法:
import torch x = torch.randn(3, 4) w = torch.randn(4, 5) y = x @ w print(y.shape) # torch.Size([3, 5])这里w的第一维是 4,正好和x的第二维匹配。如果改成w = torch.randn(3, 5),内维对不上,就会抛错。
3.3 三维批量矩阵乘法怎么理解
实际项目中经常有“一批数据”的概念。比如输入形状是(batch, seq_len, feature),权重是(batch, feature, hidden),每个样本都要做一次矩阵乘法。如果用bmm,它要求前两个 batch 维度一致,然后对每个 batch 分别做二维矩阵乘法:
batch_x = torch.randn(2, 3, 4) batch_w = torch.randn(2, 4, 5) out = torch.bmm(batch_x, batch_w) print(out.shape) # torch.Size([2, 3, 5])如果两个张量的 batch 维度不一样,比如分别是(2, 3, 4)和(3, 4, 5),bmm会失败。但matmul在这种情况下可以做广播:把第一维为 2 的batch_x和第一维为缺省或 1 的batch_w进行匹配。
x2 = torch.randn(3, 4) w_batch = torch.randn(2, 4, 5) out2 = torch.matmul(x2, w_batch) print(out2.shape) # torch.Size([2, 3, 5])这里x2被“看成一个单批量样本”,和w_batch的每个 batch 分别相乘。matmul的广播规则让代码更简洁,但也要注意:批量维度不匹配时,它可能不会立刻报错,而是按广播规则扩展,最终结果可能不是你想要的维度。
3.4 矩阵乘法报错时按什么顺序排查
矩阵乘法报错最常见的提示是size mismatch或shapes cannot be multiplied。不要一上来就改代码,先按下面顺序检查:
- 把参与运算的两个张量形状分别打印出来。
- 判断是二维还是高维。高维时先忽略 batch 维,只关注最后的二维矩阵乘是否匹配。
- 检查内维是否相等:第一个矩阵的第二个维度,是否等于第二个矩阵的第一个维度。
- 如果不相等,确认是不是需要对某一方做转置。例如
x形状(4, 3),w形状(4, 5),通常需要x.T @ w,而不是直接x @ w。 - 如果是
bmm,还要确认两个张量的 batch 维是否相等,或者是否想让matmul做广播。
4. 广播机制:从右往左对齐,缺一维就补一维
4.1 广播机制解决什么问题
上一章说两个张量形状不一致时,逐元素计算可能报错。但在很多场景下,我们希望一个标量加到张量上,或者一个向量加到矩阵的每一行上。如果每次都手动复制一份,既浪费内存,又让代码啰嗦。于是就有了广播机制。
广播不是真正把数据复制成完整矩阵,而是 PyTorch 在计算时“虚拟地”扩展维度。它的核心价值是让不同形状的张量可以做逐元素运算,同时保持代码简洁。
4.2 广播的三条规则
PyTorch 的广播规则可以概括成三条:
- 从最右边的维度开始对齐。
- 依次比较每个维度,两个维度相等,或者其中一个为 1,那么这一维可以广播。
- 如果其中一个维度缺失,就把它当成 1,再继续对齐。
看几个例子:
import torch # 标量和张量 a = torch.tensor([1.0, 2.0, 3.0]) s = torch.tensor(2.0) print(a + s) # tensor([3., 4., 5.]) # 向量和矩阵 mat = torch.randn(3, 4) vec = torch.randn(4) print((mat + vec).shape) # torch.Size([3, 4]) # 两个方向都扩展 m = torch.ones(3, 1) n = torch.ones(1, 4) print((m + n).shape) # torch.Size([3, 4])第一个例子中,s没有维度,PyTorch 会把它当成标量与张量中每个元素相加。第二个例子中,mat形状(3, 4),vec形状(4,)。从右往左看,第一维相等,都是 4,所以vec可以广播到(3, 4)。第三个例子中,m是(3, 1),n是(1, 4)。从右往左对齐:1和4,其中一个为 1,扩展为4;再左边3和1,其中一个为 1,扩展为3。最后结果就是(3, 4)。
多画几遍这个“从右往左对齐”的过程,比背结论更管用。
4.3 广播和 reshape 的结合
有时候两个张量形状看起来不匹配,但实际上只需要加一个维度就能配合。举例来说,x形状是(3, 4),v形状是(3,)。你想把v沿着“行”方向叠加?不对,广播默认是从右往左对齐的,v会优先和x的最后一维匹配。如果想让v作用于“行”,需要先把它变成(3, 1):
x = torch.ones(3, 4) v = torch.tensor([1.0, 2.0, 3.0]) print((x + v.unsqueeze(1)).shape) # torch.Size([3, 4])unsqueeze(1)是在第 1 维插入一个长度为 1 的维度,让v从(3,)变成(3, 1)。这样广播时,(3, 4)和(3, 1)从右往左对齐:4和1中有一个为 1,扩展为 4;3和3相等,保持不变。最终得到(3, 4)。
4.4 广播失效的典型情况
最常见的广播失败是类似(3,)和(4,)相加。从右往左第一个维度是 3 和 4,两个不相等而且都不为 1,所以报错。再比如(3, 4)和(2, 4),右对齐第一个维度 4 和 4 相等,第二个维度 3 和 2 不相等也不为 1,报错。
这种报错的提示已经比较清楚:The size of tensor a (3) must match the size of tensor b (2) at non-singleton dimension 1。看到这个提示,直接定位到数字,再想一下是使用unsqueeze补维、view改形状,还是根本就应该用矩阵乘法。
注意:广播机制不是万能的“自动配对”。写作时如果对输出维度不确定,建议先打印出参与运算的两个张量和结果的
shape,再继续下一步。这个习惯能帮你少踩很多隐藏坑。
4.5 该用广播的时候就别手动复制
有些初学者为了让两个张量形状一致,会写循环复制数据。比如把向量v复制成矩阵的每一行,再和矩阵相加。这样做结果当然可以,但代码更慢、更容易写错,而且内存开销更大。广播机制就是为了避免这种情况。
当然,过度依赖广播也会让代码可读性变差。如果你在一个复杂的模型里发现某个操作隐式广播,很难一眼看出哪个维度被扩展了。所以我的建议是:简单运算放心用广播;复杂运算中,尽量先用注释写出输入和输出形状,必要时用view、unsqueeze、expand把维度改清楚,再参与广播。
5. 综合示例:实现一个简化全连接层和批归一化前向过程
5.1 需求定义
现在用一个综合例子把三个规则串起来。假设输入是一个 batch 为 16、特征维度为 8 的张量:
import torch x = torch.randn(16, 8)我们要做两件事:第一,经过一个线性层,输出维度为 4;第二,对线性层输出做一次