reshape、view、transpose 和广播机制怎么选:张量形状转换实战

reshape 和 view 看起来一样,transpose 和 permute 也像兄弟,广播机制报错信息又长又看不懂。这篇用最小代码讲清它们的区别,让你不再瞎试。

学习路线:PyTorch 深度学习基础(43~58) · 第 47 课
第 47 课第 47 课张量形状转换视频封面,展示 reshape、view、transpose、permute 和广播机制的选择关系。
Learning Path

PyTorch 深度学习基础(43~58)

在前 20 篇 PyTorch 入门上继续深化,补齐张量操作、自动微分、优化器家族和训练排错能力。

学完本阶段你能做到:独立搭建一个 PyTorch 训练脚本,正确处理 Shape、device、梯度清空、优化器选择和学习率策略,并能根据 loss/accuracy 曲线诊断训练问题。

推荐读法:当前更新主线,建议跟更;每篇都要配合代码自检。

查看完整阶段 · 16 篇
Lesson Guide

这一课怎么学

先看目标,再带着问题读正文。读完后用练习确认自己真的理解了。

本节你会学到

  • 说清楚「reshape、view、transpose 和广播机制怎么选:张量形状转换实战」解决的核心问题
  • 知道它在「PyTorch 深度学习基础(43~58)」中的位置
  • 把 Tensor、DataLoader、Linear 层 这些关键词联系起来

前置知识

  • 建议先读完上一篇:Shape 是深度学习最重要的数据契约
  • 能区分输入、输出、数据和模型目标。

概念回顾

  • 【Tensor】可以先把 Tensor 理解成支持 GPU、梯度和批量计算的多维数组。 前面或后面会反复用到它。
  • 【DataLoader】DataLoader 负责把数据一批一批喂给模型,训练循环才不会手忙脚乱。 前面或后面会反复用到它。
  • 【Linear 层】Linear 层做的是矩阵乘法加偏置,是很多神经网络模块的基础零件。 前面或后面会反复用到它。

常见误区

  • 不要只记定义,要追问它解决了什么问题。
  • 不要只看 API 名字,要同时关注输入输出 shape 和训练流程。

课后练习

  • 用 3 句话向一个零基础朋友解释「reshape、view、transpose 和广播机制怎么选:张量形状转换实战」。
  • 打开概念库里的「Tensor」,补一遍它和本文的关系。
  • 读下一课「自定义 Dataset:怎样封装自己的训练数据」前,先写下你认为它会解决的问题。

自检练习与参考答案

1. 如果一段 PyTorch 代码报 device 不一致,第一步应该检查什么?

先检查模型参数、输入 Tensor、标签 Tensor 是否在同一个设备上。常见修正是把它们统一 `.to(device)`。

2. 为什么学习 PyTorch 时不能只看 API 名字?

因为训练是否正确经常取决于 shape、dtype、device 和梯度流,API 名字只能告诉你工具,不会保证数据契约正确。

3. 读完本文后,至少应该能画出哪条训练主线?

`输入数据 → forward → loss → backward → optimizer.step()`,并知道本文主题位于这条链路的哪一环。

下一步建议读「自定义 Dataset:怎样封装自己的训练数据」。

你写 x.view(2, 3) 报错,改成 x.reshape(2, 3) 就好了。你写 x.transpose(0, 1) 能跑,改成 x.permute(0, 1) 也能跑。它们到底有什么区别?什么时候用哪个?

还有广播机制——两个形状不同的张量相加,有时候自动对齐,有时候报一串看不懂的错。这篇把这些问题一次讲清。

第 47 课视频 - reshape、view、transpose 和广播机制怎么选:张量形状转换实战

概念回顾

上一篇我们讲清了 Shape 是数据契约。这一篇解决“怎么转换 Shape”。上一节你应该记住:Shape 变化有四类——增加维度、删除维度、改变维度大小、调换维度顺序。今天讲的这些函数,就是实现后两类转换的工具。


reshape vs view:能重塑但不复制

这两个函数作用几乎一样:把张量改变形状,不改变数据。

import torch

x = torch.arange(12)
print(x.shape)   # torch.Size([12])

y = x.reshape(3, 4)
z = x.view(3, 4)
print(y.shape)   # torch.Size([3, 4])
print(z.shape)   # torch.Size([3, 4])

核心区别只有一个:内存连续性要求不同。

函数 要求内存连续吗 不连续时会怎样
view 必须连续 报错 RuntimeError: view size is not compatible
reshape 不要求 自动复制一份再重塑
x = torch.randn(3, 4)

# transpose 之后内存不连续
x_t = x.transpose(0, 1)
print(x_t.is_contiguous())   # False

# view 报错
x_t.view(12)   # RuntimeError

# reshape 没事
x_t.reshape(12)   # 自动复制

实战建议:拿不准就用 reshape,它更宽容。view 性能稍好(不复制),但前提是内存连续。

-1 的用法:自动推断

两个函数都支持 -1,表示“这一维的大小你帮我算”:

x = torch.randn(32, 8, 8)

y = x.reshape(32, -1)      # [32, 64]  —— -1 自动算成 8*8=64
z = x.reshape(-1)          # [2048]    —— 展平成一维

# -1 只能出现一次
x.reshape(-1, -1)   # 报错,不知道你想怎么分

这是 Flatten 的等价写法,CNN 接全连接层时常用:

# 这两行等价
x = x.reshape(x.shape[0], -1)
x = torch.flatten(x, start_dim=1)

transpose vs permute:调换维度顺序

transpose:只换两个维度

x = torch.randn(2, 3, 4)   # [2, 3, 4]

y = x.transpose(0, 1)      # 交换第 0 和第 1 维
print(y.shape)              # [3, 2, 4]

只能一次换两个。

permute:一次性重排所有维度

x = torch.randn(2, 3, 4)   # [2, 3, 4]

y = x.permute(2, 0, 1)     # 第 2 维放最前,第 0 维居中,第 1 维最后
print(y.shape)              # [4, 2, 3]

permute 的参数是新顺序中,每个位置放原来的第几维。

最经典的图像场景

图像库读出来的图片通常是 [H, W, C](高、宽、通道),但 PyTorch CNN 要 [C, H, W]

# [H, W, C] → [C, H, W]
img = torch.randn(32, 32, 3)        # 假设是 [H, W, C]
img_chw = img.permute(2, 0, 1)      # [3, 32, 32]
print(img_chw.shape)                # torch.Size([3, 32, 32])

transpose 之后想用 view?先 contiguous

transposepermute 返回的是视图,不复制数据,所以内存通常不连续。想在这之后用 view,得先 contiguous()

x = torch.randn(3, 4)
x_t = x.transpose(0, 1)       # 不连续

# x_t.view(12)                # 报错
x_t = x_t.contiguous()        # 复制成连续内存
x_t.view(12)                  # 现在能用了

或者直接用 reshape,它不挑食:

x_t = x.transpose(0, 1)
y = x_t.reshape(12)           # 直接能用,内部自动处理

squeeze 和 unsqueeze:增删大小为 1 的维度

unsqueeze:加一个大小为 1 的维度

x = torch.randn(10)           # [10]
y = x.unsqueeze(0)            # [1, 10]  —— 在第 0 维加
z = x.unsqueeze(1)            # [10, 1]  —— 在第 1 维加

最常见的用途:单张图片加 batch 维度。

img = torch.randn(3, 32, 32)        # [C, H, W]
img_batch = img.unsqueeze(0)        # [1, C, H, W]  —— 模型需要 batch 维

squeeze:删掉大小为 1 的维度

x = torch.randn(1, 32, 1)     # [1, 32, 1]
y = x.squeeze()               # [32]  —— 删掉所有大小为 1 的维度
z = x.squeeze(0)              # [32, 1] —— 只删第 0 维

危险操作:不带参数的 squeeze() 会删掉所有大小为 1 的维度,可能误删 batch 维:

x = torch.randn(1, 1, 32)     # batch=1, channel=1, length=32
x.squeeze()                   # [32] —— batch 和 channel 都删了!可能不是你想要的

# 推荐指定维度
x.squeeze(1)                  # [1, 32] —— 只删 channel 维

广播机制:形状不同也能运算

当两个形状不同的张量做 +-*/ 时,PyTorch 会自动“拉伸”某些维度让它们对齐。这就是广播。

核心法则:右对齐检查法

把两个形状从右往左对齐,每一维检查是否满足以下三个条件之一:

  1. 两个数字相等
  2. 其中一个是 1(会被拉伸)
  3. 其中一个不存在(左边空出来,等同 1)

三个条件都不满足——报错。

# 案例 1:成功
a = torch.ones(3, 1)     # [3, 1]
b = torch.ones(1, 2)     # [1, 2]
c = a + b                # [3, 2] —— 两个都拉伸

分析:

  • 最后一维:12 → 满足条件 2,a 拉伸成 2
  • 第一维:31 → 满足条件 2,b 拉伸成 3
# 案例 2:失败
a = torch.ones(4, 3)     # [4, 3]
b = torch.ones(3, 3)     # [3, 3]
c = a + b                # 报错!

分析:

  • 最后一维:33 → 满足条件 1
  • 第一维:43 → 三个条件都不满足,报错

深度学习里的两个经典场景

场景 1:偏置相加

全连接层输出 [B, C],偏置是 [C]。广播自动把偏置拉伸成 [B, C]

logits = torch.randn(32, 10)   # [B, C]
bias = torch.randn(10)          # [C]
result = logits + bias          # [32, 10] —— bias 广播成 [32, 10]

场景 2:归一化

图像 [B, C, H, W] 减去每个通道的均值 [C, 1, 1]

features = torch.randn(32, 3, 224, 224)
mean = torch.randn(3, 1, 1)          # [C, 1, 1]
normalized = features - mean          # [32, 3, 224, 224] —— 自动广播

速查表:什么时候用什么

需求 用什么 例子
改变形状(总元素数不变) reshape [B,C,H,W] → [B, C*H*W]
改变形状,且内存确定连续 view 同上,性能稍好
交换两个维度 transpose [B,T,D] → [T,B,D]
重排多个维度 permute [B,H,W,C] → [B,C,H,W]
加一个大小为 1 的维度 unsqueeze [F] → [1, F]
删掉大小为 1 的维度 squeeze(dim) [1, B] → [B]
不同形状的元素级运算 广播(自动) [B,C] + [C]

三个高频错误

错误 1:transpose 后用 view 报错

x = torch.randn(3, 4)
x_t = x.transpose(0, 1)
x_t.view(12)   # RuntimeError: view size is not compatible

修复:用 reshape 或先 contiguous()

错误 2:squeeze 误删 batch 维

x = torch.randn(1, 10)   # batch=1
y = x.squeeze()           # [10] —— batch 维没了!
model(y)                  # 报错:模型要 [B, F],你给了 [F]

修复:指定维度 x.squeeze(1),或者根本别 squeeze batch 维。

错误 3:广播静默错误

logits = torch.randn(32, 10)   # [B, C]
y = torch.randn(32, 10)         # 本意是 [B],却写成了 [B, 10]
loss = logits + y               # 不报错!但语义完全错了

广播不会报错,但结果可能完全不是你想要的。养成习惯:运算前 print(x.shape) 确认。


课后练习

练习 1:把 [2, 3, 4] 的张量变成 [6, 4],写出三种不同的写法。

练习 2:下面代码能跑吗?为什么?结果 Shape 是多少?

a = torch.randn(2, 1, 3)
b = torch.randn(4, 3)
c = a + b

练习 3:CNN 输出 [32, 64, 7, 7],要接 Linear(3136, 10)。写一行代码把 CNN 输出转成 Linear 能接收的形状。

参考答案 / 自检思路

练习 1

x = torch.randn(2, 3, 4)
# 写法一
x.reshape(6, 4)
# 写法二
x.view(6, 4)        # 原始内存连续时可用
# 写法三
x.permute(0, 1, 2).contiguous().view(6, 4)  # 或 x.flatten(0, 1)
# 也可以
x.reshape(6, -1)

练习 2:能跑。右对齐分析:

  • 最后一维:33 → 相等,通过
  • 中间维:14 → 1 会广播成 4
  • 第一维:2不存在 → 缺位等同 1,广播成 2
  • 结果 Shape:[2, 4, 3]

练习 3

x = torch.randn(32, 64, 7, 7)
x = x.reshape(32, -1)   # 或 x.reshape(x.shape[0], -1) 或 torch.flatten(x, 1)
# 现在 x.shape = [32, 3136],可以接 Linear(3136, 10)

验证:64 * 7 * 7 = 3136,和 Linear 的 in_features 一致。


核心要点小结

  • reshapeview 宽容——内存不连续也能用,拿不准就 reshape
  • transpose 换两个维度,permute 重排所有维度
  • transpose/permute 后内存不连续,想用 view 要先 contiguous
  • squeeze 别乱用不带参数的版本,容易误删 batch 维
  • 广播是自动的,但“不报错”不等于“结果对”,运算前 print shape
  • 速查表:改形状 reshape、换维度 permute、加维度 unsqueeze、删维度 squeeze(dim)

下一篇我们离开形状转换,回到数据加载——自定义 Dataset 怎么封装你自己的训练数据。

Discussion

留言讨论