1.3 张量操作与形状变换 本节摘要:训练流水线里一半的报错是形状报错。本节给出形状变换四件套(reshape/view、unsqueeze、permute、expand)的适用场景,讲透广播机制的三条规则,并用一张图把"数据在维度间的流动"画出来。 变形是日常,不是特殊操作 装备箱认过了单个张量,现在看它们怎么组合作业。为什么变形如此高频?因为各组件对形状的期待不一样:卷积层要四维输入,全连接层要二维,损失函数有时要压掉一维。数据从磁盘到 loss,途中被变形十几次是常态。形状意识不是加分项,是及格线。 先看一条总原则,能少走一半弯路:view 和 reshape 不搬运数据,只改"解读方式"。张量的数据在内存里是一维排布的,形状只是贴在外面的解读标签。
本节摘要:训练流水线里一半的报错是形状报错。本节给出形状变换四件套(reshape/view、unsqueeze、permute、expand)的适用场景,讲透广播机制的三条规则,并用一张图把"数据在维度间的流动"画出来。
装备箱认过了单个张量,现在看它们怎么组合作业。为什么变形如此高频?因为各组件对形状的期待不一样:卷积层要四维输入,全连接层要二维,损失函数有时要压掉一维。数据从磁盘到 loss,途中被变形十几次是常态。形状意识不是加分项,是及格线。
先看一条总原则,能少走一半弯路:view 和 reshape 不搬运数据,只改"解读方式"。张量的数据在内存里是一维排布的,形状只是贴在外面的解读标签。这解释了为什么 view 快、也解释了它什么时候会失败。

右下那块"事故现场"不是吓唬人,四个场景在本册后面的代码里会出现三个,届时你会认出它们。
import torch img = torch.randn(64, 1, 8, 8) # 一个batch的灰度小图:批、通道、高、宽 print("原始形状:", img.shape) # view/reshape:展平成全连接层要吃的二维。-1 表示"这一维你替我算" flat = img.view(64, -1) print("展平后:", flat.shape) # 64 x 64 # unsqueeze:在指定位置补一个长度为1的维度 x = torch.randn(64, 8, 8) x_batched = x.unsqueeze(1) # 在第1维插入通道维 print("补通道后:", x_batched.shape) # 64 x 1 x 8 x 8 # permute:调换维度顺序(注意与 transpose 的区别:permute 可一次调任意多个维) feat = torch.randn(64, 8, 8, 3) # 假设这是"通道在最后"的布局 feat_torch = feat.permute(0, 3, 1, 2) print("调序后:", feat_torch.shape) # 64 x 3 x 8 x 8
输出:
原始形状: torch.Size([64, 1, 8, 8]) 展平后: torch.Size([64, 64]) 补通道后: torch.Size([64, 1, 8, 8]) 调序后: torch.Size([64, 3, 8, 8])
三个细节值得留意。第一,view(64, -1) 里的 -1 是让 PyTorch 自动推算该维大小,只能有一个 -1。第二,unsqueeze 补的是长度为 1 的维度,它不复制任何数据。第三,permute 之后张量通常不再连续(内存里数据还是老顺序,标签却换了),此时再 view 会报错——报错信息里出现 contiguous 就是被它撞上了,跟一句 .contiguous() 或改用 reshape 即可。
背景:你想把 permute 过的特征图展平,随手写了 feat_torch.view(64, -1),运行报错,信息里带着 view size is not compatible ... use .reshape 字样。
操作:复现并对比三种写法。
feat = torch.randn(64, 8, 8, 3) p = feat.permute(0, 3, 1, 2) # permute 后内存布局与标签不再一致 try: p.view(64, -1) # 大概率报错 except RuntimeError as e: print("view 报错(截断):", str(e)[:60]) ok1 = p.contiguous().view(64, -1) # 修法一:先恢复连续再 view ok2 = p.reshape(64, -1) # 修法二:reshape 自动处理(必要时复制) print("两种修法结果一致:", torch.equal(ok1, ok2), ok2.shape)
输出:
view 报错(截断): view size is not compatible with input tensor's s 两种修法结果一致: True torch.Size([64, 192])
解读:view 只肯重贴标签、拒绝搬数据,遇到"标签与内存不匹配"就罢工;reshape 的策略是能不搬就不搬、必要时复制一份。两者结果数值相同,差别只在是否可能产生数据拷贝。性能敏感的热路径上优先 view,随手的胶水代码用 reshape 更省心。
变式:把 permute 换成 feat.view(64, -1) 直接对原张量操作,根本不会报错——因为原张量是连续的。这提醒我们:报错的根源不是 view 本身,而是之前的维度操作改变了连续性。
两个形状不同的张量做运算时,PyTorch 按三条规则从后往前对齐维度:维度数不足就在前面补 1;长度为 1 的维度自动拉伸复制;其余维度必须相等或为空。这套机制省掉了大量手写 expand 的样板代码,但也带来"形状悄悄对上、语义悄悄错了"的隐性风险。
# 合法广播:输出层偏置按行加到每个样本上 scores = torch.randn(64, 10) # 每个样本10个类别得分 bias = torch.randn(10) # 每个类别一个偏置 out = scores + bias # bias 被广播成 64x10 print("广播结果:", out.shape) # 危险广播:形状能对上,但语义错了 a = torch.randn(64, 10) b = torch.randn(64, 1) c = a + b # 合法:b 沿第1维复制 print("对齐方向要注意:", c.shape) # 若本意是"每行加一个标量",这里就加错了方向
输出:
广播结果: torch.Size([64, 10]) 对齐方向要注意: torch.Size([64, 10])
解读:第二个例子里 a + b 完全合法,但 b 是沿着哪一维复制、是不是你想要的语义,机器不会替你判断。规则很简单:凡是依赖广播的运算,写完立刻 print 一次小尺寸结果人工核对,比事后调试省一小时。
变式:想强制广播失败来暴露语义错误,可以对 b 先做 b.squeeze(1) 或明确 b.expand_as(a),让广播意图显式化——团队代码里推荐后者,因为意图写在了代码里。
装备库清点完毕。第 2 章开始,这些张量将被装进"层"里,组装成真正的远征队——nn.Module 登场。