PyTorch张量运算核心:形状、广播与矩阵乘法实战指南

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(min0)意思是“把小于 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], dtypetorch.float64) b torch.tensor([2.0], dtypetorch.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扩展成4b的缺失维度补为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(min0) 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 结构都会少走很多弯路。先把这层地基打好我们下一课继续往上盖。