PyTorch张量运算核心:形状、广播与矩阵乘法实战指南
很多人在刚接触 PyTorch 时,会经历一个看起来不起眼、实际上非常关键的分水岭:装好环境、跑通print(torch.__version__)之后,跟着教程敲几行张量运算,加加减减都能正常输出,感觉很容易。可一旦开始写真实模型,遇到高维矩阵乘法,形状对不上了;遇到广播机制,结果莫名其妙变成一个更大的矩阵;甚至同一段代码,在 CPU 上可以运行,换到 GPU 上就开始报错。你会发现,前面那些“太简单了”的逐元素计算、矩阵乘法、广播机制,其实藏着整套框架的心智模型。这也是我这篇 PyTorch 第 2 课最想解决的问题:不是帮你背 API,而是把张量运算背后那套规则真正拆开,让你之后看到任何一段模型代码,都能在脑子里画出“形状是怎么流动的”。
先说一个判断:掌握张量运算规则,90% 的功夫在形状,10% 才在 API。torch.add、torch.matmul、torch.mul这些函数你一时想不起来都可以查文档,但如果你不知道参与运算的两个张量分别是什么形状、结果应该长什么样、维度对齐的时候谁在跟谁对齐,那代码就只能在“试错 - 看报错 - 改 shape”里打转。本文会把这一课拆成六个部分:先帮你建立一套形状视角,再分别讲透逐元素计算、矩阵乘法、广播机制,最后用一个小实战和一条排查链路,把这些规则真正落到代码里。
1. 先建立一张“形状流程图”,再谈任何运算
1.1 张量形状不只是描述,而是一份尺寸合同
很多初学者会把shape当作一个“输出时打印出来的信息”,比如torch.Size([2, 3]),看一眼就过了。但更准确的理解是:shape 是一份合同,它约定了数据如何排布,也约定了当前张量能跟谁运算、不能跟谁运算。
比如torch.arange(6).reshape(2, 3)会得到一个形状为(2, 3)的张量。在内存里,它仍然是 6 个连续排布的数字,但是 PyTorch 会按照“2 行 3 列”的方式去解释这段数据。后续的相加、相乘、转置,全都建立在这一份“解释”之上。你改变了 shape,等于改变了对同一段数据的切分方式,运算规则也随之改变。
所以我在看别人代码或自己写代码时,第一件要做的事永远是:把输入、权重、偏置、输出这几件事的 shape 列出来。列完之后再动手写运算,效率会高很多。
1.2 用“从右往左对齐”的眼光看所有张量运算
这里给出一个全篇最重要的方法:遇到任何张量运算,先把参与运算的每一个张量的 shape 写出来,然后从最后一个维度开始,逐维对齐。其实无论逐元素运算、矩阵乘法还是广播机制,底层都离不开“最后维对齐”这件事。
举个例子。假设a的形状是(3, 4),b的形状是(4,)。如果执行a + b,你会发现它不会报错,因为b的(4,)会和a的最后一维(4,)对齐,然后b相当于在行方向上被“复制扩展”到(3, 4),最终得到一个(3, 4)的结果。这就是广播机制的雏形。
但如果执行a @ b(矩阵乘法),规则就不一样了。矩阵乘法要求a的最后一维等于b的倒数第二维。a是(3, 4),b是(4,),矩阵乘法会把b当作(4, 1)来参与运算,最后得到(3, 1),再压缩成(3,)。你看,同样两个张量,一个用逐元素加法,一个用矩阵乘法,判断的起点其实都是从最后一个维度开始。所以请记住:看到任何运算,先看 shape,再从右往左对齐。
import torch a = torch.randn(3, 4) b = torch.randn(4) # 逐元素加法:形状可广播,结果为 (3, 4) c = a + b print(c.shape) # torch.Size([3, 4]) # 矩阵乘法:b 被当作 (4, 1),结果为 (3,) d = a @ b print(d.shape) # torch.Size([3])这个“从右往左对齐”的习惯,会在接下来每一节里反复出现。
2. 逐元素计算:它比看起来更容易踩坑
2.1 逐元素的本质是“同一位置的数据,同一套逻辑”
逐元素计算,指两个张量在相同位置上的元素分别进行运算。+、-、*、/、比较运算、torch.where、torch.clamp,甚至大部分激活函数,本质上都是逐元素操作。
初看很容易:两个形状完全相同的张量,对应位置相加,结果还是同一个形状。这句话没错,但真正的问题在于,“形状完全相同”只是最简单的情况。当两个张量形状不同时,PyTorch 可能不会报错,而是自动启用广播;启用广播后,运算仍然是逐元素的,只是参与运算的元素范围被“扩展”了。
这就引出一个非常关键的认知:逐元素计算并不关心数据的语义,它只关心位置对应关系。位置一旦错位,结果不会报错,但可能全错。比如(3, 1)和(1, 3)相加,结果是(3, 3),如果你原本以为是在对齐(3,)与(3,),就会得到完全不符合预期的矩阵。
x = torch.ones(3, 1) y = torch.ones(1, 3) z = x + y print(z.shape) # torch.Size([3, 3])这种“运算合法但语义错误”的情况,是逐元素计算中最隐蔽的坑。
2.2 很多看似高级的层,底层都是逐元素运算
理解逐元素的重要性,不是因为你以后会天天手写a + b,而是因为神经网络里大量的算子,本质上就是逐元素逻辑。
举个例子:ReLU 激活函数,在旧代码里经常是x.clamp(min=0),意思是“把小于 0 的元素变成 0,其他元素不变”。这不就是逐元素计算吗?还有批量归一化里的缩放和平移:(x - mean) / sqrt(var + eps) * gamma + beta,虽然从公式上看比较复杂,但落到张量上,依然是逐元素操作,因为mean、var、gamma、beta通常都沿着通道维度做了广播,最终在每个位置上独立运算。
理解了这一点,你会更容易明白为什么 GPU 对深度学习这么重要:GPU 最擅长的就是大量相互独立的逐元素计算并行执行。一个 1 万乘以 1 万的张量做逐元素相乘,在 GPU 上可能只是一瞬间的事,因为它不需要串行等待,每个位置都能同时算。
2.3 两个新手最容易忽略的工程细节
逐元素计算入门容易,工程上却有几个容易忽略的细节。
第一,in-place 操作要非常谨慎。x.add_(1)是在原张量上直接修改,如果这个张量参与了自动求导,可能导致梯度计算出错。你可能会看到RuntimeError: a leaf Variable that requires grad is being used in an in-place operation这类报错。工程上的建议是:除非明确需要省内存,否则尽量写x = x + 1而不是x.add_(1)。
第二,dtype 和精度问题。默认的浮点类型是torch.float32,如果两个张量一个float32一个float64,直接相加可能报类型不匹配。而在做累加时,大量float32数字相加可能会有精度损失。实际项目里,这种情况经常出现在 loss 累加、梯度累积或者归一化统计量计算中。遇到精度敏感的场景,要么显式转dtype,要么考虑中间结果用更高精度。
a = torch.tensor([1.0], dtype=torch.float64) b = torch.tensor([2.0], dtype=torch.float32) # 下面的代码会报 dtype 不匹配 # c = a + b # 可以先统一类型 c = a + b.to(torch.float64)3. 矩阵乘法:忘掉 API,记住维度的握手
3.1 一次典型报错的完整拆解
很多人在模型代码里看到矩阵乘法时最头疼的报错是:
RuntimeError: mat1 and mat2 shapes cannot be multiplied (2x4 and 3x5)这句话其实已经把原因说得很清楚了:左边矩阵的列数(4)不等于右边矩阵的行数(3)。但为什么代码会写出这种形状错配?通常是因为前面的 reshape、transpose、squeeze 让数据顺序发生了变化,或者你根本不理解某一步之后张量变成了什么形状。
遇到这种报错,第一步不是在代码里瞎改transpose,而是回到数据流里,把每一步的 shape 都打印出来。我见过不少同学,看着报错信息说“是不是要换torch.bmm”,其实问题根本不在 API,而是两个张量本来就不满足矩阵乘法的维度要求。
3.2 一维、二维、高维,底层只有一套规则
PyTorch 里的矩阵乘法常用三种写法:torch.mm、torch.matmul和@运算符。torch.mm只支持二维矩阵;torch.matmul和@更通用,支持高维张量。列成一个表会更清晰。
| 情况 | 示例 | 规则 | 结果形状 |
|---|---|---|---|
| 一维 × 一维 | (D,) @ (D,) | 向量点积 | 标量,shape 为() |
| 一维 × 二维 | (D,) @ (D, H) | 向量作为(1, D)参与运算 | 结果先得到(1, H),实际返回(H,) |
| 二维 × 一维 | (N, D) @ (D,) | 向量作为(D, 1)参与运算 | 结果先得到(N, 1),实际返回(N,) |
| 二维 × 二维 | (N, D) @ (D, H) | 矩阵乘法 | (N, H) |
| 高维 × 高维 | (B, N, D) @ (B, D, H) | 后两维做矩阵乘法,前面的 batch 维必须逐维对齐 | (B, N, H) |
所以你应该发现,无论是一维、二维还是更高维,底层规则只有一条:参与矩阵乘法的两个张量,从最后两个维度看,必须满足“前一个的最后维 == 后一个的倒数第二维”。前面所有维度当作 batch 维处理,要求逐维一致,或者满足广播条件。
举个实际例子:
x = torch.randn(4, 8) # 输入:4 个样本,每个样本 8 维特征 w = torch.randn(8, 16) # 权重:将 8 维输入映射到 16 维隐藏层 out = x @ w # 结果:4 个样本,每个样本 16 维 print(out.shape) # torch.Size([4, 16])这段代码背后的 shape 变化是:(4, 8) @ (8, 16) -> (4, 16)。中间那对 8 被“吃掉”了。
3.3 转置与维度交换:矩阵乘法最容易出错的另一半
矩阵乘法常见的另一个错误来源是转置。在做y = x @ w时,经常需要判断是w还是w.T。这里的记忆口诀是:最终结果的最后一维,来自第二个矩阵的最后一维;最终结果的前面维度,来自第一个矩阵的前面维度。
比如在神经网络中,一个常见的操作是把(B, T, D)的特征和(D, H)的权重相乘,得到(B, T, H)。这里权重(D, H)意味着把输入的特征从D维映射到H维,不需要转置。但是如果你面对的是(B, D, T),想得到(B, T, H),就得先把形态调整成(B, T, D),这时候transpose或者permute就会登场。
一个通用的实操建议是:在编写任何涉及矩阵乘法的代码前,先写一行注释,把 shape 的变换过程写出来。例如:
# x: (B, T, D) # w: (D, H) # out: (B, T, H) out = x @ w这一步看起来简单,却能省下大量排查时间。
4. 广播机制:它才是 PyTorch 张量运算的灵魂
4.1 广播不是在复制数据,而是在对齐维度
广播机制是 PyTorch 张量运算里最灵活、也最需要正确心智模型的部分。很多人理解成“把小张量复制成和大张量一样的形状,然后再运算”,这个说法在直觉上没错,但它会让你误以为内存会被展开、速度会变慢。实际上,PyTorch 在执行广播时,通常会通过底层的 stride 机制和向量化计算来实现逻辑上的扩展,而不是真的把数据复制成完整的多份。
正确的心智模型是:广播是维度对齐的一种规则。
规则有三条:
- 从最后一个维度开始,逐维向前比较。
- 如果两个维度相等,则保留该维度。
- 如果两个维度不相等,但其中一个维度为 1,则这个维度可以被扩展为另一个张量的维度;如果一个张量没有某个维度,则视为 1。
如果两个维度既不相等,也没有一个是 1,就会报错。
a = torch.randn(3, 1) # 第二维是 1 b = torch.randn(4) # shape: (4,) c = a + b # 对齐后:a 变成 (3, 4),b 变成 (3, 4),结果为 (3, 4) print(c.shape) # torch.Size([3, 4])这个例子里,a的形状(3, 1)和b的形状(4,)被视为(1, 4),在维度对齐时,a的1扩展成4,b的缺失维度补为1再扩展成3,最终结果是(3, 4)。
4.2 广播省代码,也会隐藏最危险的结果错误
广播机制的价值在于代码简洁。一个最经典的例子就是偏置项的加法:x的形状是(B, D),偏置b的形状是(D,)。如果没有广播,你需要先把breshape 成(1, D),然后复制B份变成(B, D),再加到x上。有了广播,直接x + b就行。
但风险也随之而来。当两个张量的形状看起来很相似,实际上并不直接匹配时,广播可能给你“意外的成功”。比如x的形状是(3,),y的形状是(3, 1),如果你在代码里把它们相加,结果会是(3, 3),而且不会报任何错。这种“没有异常、但结果不对”的 bug 是广播机制里最危险的类型,因为它很难被一开始的 try-except 拦截,只有当你盯着输出矩阵看很久,才会发现维度默默扩大了。
建议:在关键运算后加一行 shape 断言,例如
assert result.shape == expected_shape,让“隐式广播”变成“显式检查”。这在写模型和训练循环时尤其有用。
4.3 用三步判断法预测广播结果
我一般会用一个三步法来判断广播后的结果形状:
- 写出所有参与张量的 shape。
- 从最后一个维度开始,逐维对齐;缺失维度补 1,维度为 1 则可以扩展。
- 对齐后,每一维取两个张量中的最大值,作为结果形状。
举例验证:
(3, 4)和(4,):对齐后是(3, 4)与(1, 4),结果为(3, 4)。(3, 1)和(1, 4):对齐后是(3, 1)与(1, 4),结果为(3, 4)。(3, 4)和(3, 1):对齐后是(3, 4)与(3, 1),结果为(3, 4)。(3, 4)和(4, 3):最后一位是4和3,不相等,且都不是 1,直接报错。
用代码验证一下:
import torch def check_broadcast(a_shape, b_shape): a = torch.randn(a_shape) b = torch.randn(b_shape) try: c = a + b print(f"{a_shape} + {b_shape} -> {c.shape}") except RuntimeError as e: print(f"{a_shape} + {b_shape} -> 广播失败: {e}") check_broadcast((3, 4), (4,)) check_broadcast((3, 1), (1, 4)) check_broadcast((3, 4), (4, 3))输出会很清楚:前两个成功,并且结果形状符合预期;第三个会在运行时直接抛出维度不匹配的异常。
5. 一个不依赖 nn.Module 的小实战:手写两层前向传播
5.1 先定好网络形状,再写张量运算
很多教程在介绍完张量运算后,会直接跳到nn.Linear、nn.Sequential。但我的建议是,如果你想真正理解张量运算规则,先不要急着用高级封装,而是用最原始的@、加法、逐元素函数手写一个两层网络的前向传播。
假设输入X的形状是(N, D),其中N是样本数,D是特征数。我们要实现一个隐藏层,隐藏单元数为H,输出类别数为C。那么权重和偏置的设计是:
W1:(D, H)b1:(H,)W2:(H, C)b2:(C,)
这里有个值得注意的细节:为什么偏置是(H,)而不是(1, H)?因为(H,)可以直接通过广播加到(N, H)上,结果仍然是(N, H)。如果你写的是(1, H),广播规则也能工作,但多了一个不必要的维度。从直觉上说,把偏置理解成“每个隐藏单元有一个标量偏移量”,而不是“一行向量”,更符合后面的梯度计算习惯。
5.2 关键代码:每行都标出 shape 变化
import torch def relu(x): # 逐元素计算:x: (N, H) -> (N, H) return x.clamp(min=0) N, D, H, C = 8, 16, 32, 10 x = torch.randn(N, D) # x: (8, 16) w1 = torch.randn(D, H) # w1: (16, 32) b1 = torch.randn(H) # b1: (32,) # 线性层 1 z1 = x @ w1 + b1 # (8, 16) @ (16, 32) -> (8, 32); 广播 + b1 -> (8, 32) a1 = relu(z1) # a1: (8, 32) w2 = torch.randn(H, C) # w2: (32, 10) b2 = torch.randn(C) # b2: (10,) # 线性层 2(输出层) logits = a1 @ w2 + b2 # (8, 32) @ (32, 10) -> (8, 10); 广播 + b2 -> (8, 10) print(logits.shape) # torch.Size([8, 10])可以看到,每一行代码都在做同一件事:先确定当前输入的 shape,再选择匹配的权重和偏置,然后通过@和+得到下一层输出。当你把 shape 注释写在每一步旁边时,整个网络结构就变得非常直观。
5.3 用打印和断言养成“每步验证”的习惯
上面这个例子能跑通,但如果有一天你写了一个更复杂的网络,某个中间维度的 shape 错了,报错信息可能离真正的问题很远。所以我建议在每一步后面加一个可选的断言,或者干脆自定义一个简单的打印函数:
def trace(label, tensor): print(f"{label}: {tuple(tensor.shape)}") return tensor x = torch.randn(8, 16) w1 = torch.randn(16, 32) b1 = torch.randn(32) z1 = trace("z1", x @ w1 + b1) # z1: (8, 32) a1 = trace("a1", relu(z1)) # a1: (8, 32) w2 = torch.randn(32, 10) b2 = torch.randn(10) logits = trace("logits", a1 @ w2 + b2) # logits: (8, 10)一旦某个 shape 不是预期值,你立刻能定位到是哪一步出了问题。这种“每步打印”的习惯,在写 Transformer、RNN 等复杂结构时,价值会被放大到和模型正确性直接相关。
6. 张量运算报错与隐性问题排查链路
6.1 常见报错与真正根因对照
| 报错信息 | 常见根因 | 排查侧重 |
|---|---|---|
mat1 and mat2 shapes cannot be multiplied | 矩阵乘法前后维度不匹配 | 检查参与运算的两个张量 shape,确认前一个最后维 == 后一个倒数第二维 |
The size of tensor a must match the size of tensor b | 逐元素运算形状不匹配,且不满足广播条件 | 从最后一维开始逐维对比,找出不匹配的维度 |
Sizes of tensors must match except in dimension 1 | 常见于 cat、stack 等拼接操作 | 确认拼接维之外的维度是否完全一致 |
a leaf Variable that requires grad is being used in an in-place operation | 对需要梯度的叶子张量做了原地修改 | 检查是否有add_、mul_、copy_等 in-place 操作 |
| 没有报错,但结果 shape 比预期大很多 | 广播让维度意外扩展,导致结果不符合语义 | 复算广播结果,在关键运算后加断言 |
这些报错里,前四个都是显式的,最后一种却是隐藏的。显式报错并不可怕,因为它会打断你,逼你去查;真正要警惕的是“成功运行但结果错误”的隐性问题。
6.2 一条可复用的五步排查顺序
当张量运算出现异常,我推荐按下面的顺序排查,而不是直接去改代码。
- 看现象:先确认是报错、卡住、输出 shape 异常,还是结果数值不对。不同的现象指向不同层次的问题。
- 核对 shape:把出错的运算写成注释或打印出来,列出所有参与张量的 shape。这一步能解决大部分问题。
- 核对维度顺序:确认是否需要
transpose、permute、reshape或squeeze。尤其是从图像、文本、序列数据进入网络时,维度顺序最容易混乱。 - 用极小例子复现:如果 shape 看起来没问题,但结果不对,就构造一个非常小的随机张量,手动算一遍预期结果,再与代码输出对比。比如两个
(2, 2)矩阵相乘,手算一下结果,立刻能发现是广播还是转置的问题。 - 给关键运算加断言:在关键步骤后加
assert result.shape == expected_shape,避免隐性错误再次静默通过。
这套顺序,本质上是从“现象”逐步深入到“数据的形状和语义”再到“代码的工程防护”。
6.3 版本变化也会引发“运算没写错但就是要改”的情况
还有一种情况需要提醒:你写的张量运算本身没问题,但 PyTorch 版本升级后,某些 API 的默认行为发生变化。在较新的 PyTorch 版本(例如社区里讨论比较多的 2.6 系列)中,torch.load的weights_only参数默认值就发生了一处调整,导致加载旧模型时可能出现提示。这不是张量运算写错了,而是框架在安全性和兼容性之间做了新的取舍。
遇到这类问题,第一步不是改[0]或者强行绕开,而是去查对应版本的 release notes 或迁移说明,看清楚行为变更的背景,再决定怎么改。这也提醒我们:在本地搭建 PyTorch 环境时,锁好版本、记录依赖,不只是一件安装的事,它会直接影响你在后续训练和部署时能否保持行为一致。无论你用的是 CPU 版还是 GPU 版,这个原则都成立。
说到环境安装,可能有的同学在前一步已经被“下载很慢”折磨过。这件事本课不展开,但可以给一个方向:下载慢通常不是框架本身慢,而是默认源或网络链路的问题,换一个合适的镜像源往往就能解决。安装完、确定版本之后,再回到张量运算时,你会发现那些规则不会因为版本变化而失效。
回到本课最核心的经验:张量运算的规则并不复杂,复杂的是让“眼睛看到的 shape”和“脑子里理解的 shape”保持一致。最笨但最有效的方法是:动手前先画,画完再跑;跑一步,打印一次 shape,直到每一步都对上。当你真正养成“先画 shape、再写运算、最后验证输出”的习惯,后面的反向传播、卷积、Transformer 结构,都会少走很多弯路。先把这层地基打好,我们下一课继续往上盖。
