Skip to content

第 5 章 形状操作

学习目标

  • 掌握 reshapeview,理解两者的区别
  • 掌握转置 T / transpose / permute
  • 掌握 squeezeunsqueeze 增删维度
  • 掌握拼接 cat 与堆叠 stack、拆分 split / chunk
  • 理解「连续性(contiguous)」与 flatten

5.1 reshape:改变形状,不改数据

reshape(重塑)按新形状重新解释同一批数据。元素总数必须不变:

python
import torch

a = torch.arange(6)          # [0, 1, 2, 3, 4, 5]
print(a.reshape(2, 3))       # 2 行 3 列
print(a.reshape(3, -1))      # -1 自动推断:3 行 → 每行 2 个

输出:

tensor([[0, 1, 2],
        [3, 4, 5]])
tensor([[0, 1],
        [2, 3],
        [4, 5]])

-1 是「让 PyTorch 根据元素总数自动算」的占位符,一维形状中只能出现一个 -1。数据顺序不变:无论形状怎么变,元素仍按 0, 1, 2, ... 的行优先顺序排列。

5.2 view:reshape 的「零拷贝」版本

viewreshape 表面一样:

python
import torch

a = torch.arange(6)
print(a.view(2, 3))

输出:

tensor([[0, 1, 2],
        [3, 4, 5]])

区别在于:view 不复制数据,只是换一种方式「看」同一块内存,所以要求原张量在内存里是连续(contiguous)的。reshape 在必要时会自动复制,更省心,是新手首选。

什么情况会不连续?最常见的是转置之后:

python
import torch

nt = torch.arange(6).reshape(2, 3).T
print("是否连续:", nt.is_contiguous())

输出:

是否连续: False

对不连续张量调用 view 会报错:

>>> nt.view(-1)
Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
RuntimeError: view size is not compatible with input tensor's size and stride (at least one dimension spans across two contiguous subspaces). Use .reshape(...) instead.

报错最后一句直接给了建议:Use .reshape(...) instead.reshape 可以安全处理不连续张量:

python
print(nt.reshape(-1))

输出:

tensor([0, 3, 1, 4, 2, 5])

注意顺序:转置后 0 3 / 1 4 / 2 5 逐行展开,所以是 [0, 3, 1, 4, 2, 5]

5.3 转置与 permute

转置(transpose)交换两个维度。二维用 .Ttranspose(0, 1):

python
import torch

m = torch.arange(6).reshape(2, 3)
print(m.T)
print(m.transpose(0, 1))

输出(两种写法结果相同):

tensor([[0, 3],
        [1, 4],
        [2, 5]])
tensor([[0, 3],
        [1, 4],
        [2, 5]])

三维及以上的任意维度重排用 permute(排列),参数是新维度的顺序:

python
import torch

t = torch.arange(24).reshape(2, 3, 4)   # 形状 (2, 3, 4)
print(t.permute(2, 0, 1).shape)          # 新顺序:第 2、0、1 维

输出:

torch.Size([4, 2, 3])

permute(2, 0, 1) 把原来的第 2 维(长度 4)放到最前面。转置、permute 都不复制数据,返回的视图通常不连续。

5.4 squeeze 与 unsqueeze:删掉和加上「长度为 1」的维度

长度为 1 的维度不影响数据,但影响形状匹配(第 4 章广播)和网络层的输入要求。squeeze(压缩)删除所有长度为 1 的维度,或指定删除某一个:

python
import torch

x = torch.zeros(1, 3, 1)
print(x.squeeze().shape)     # 全部删除 → (3,)
print(x.squeeze(0).shape)    # 只删第 0 维 → (3, 1)
print(x.squeeze(2).shape)    # 只删第 2 维 → (1, 3)

输出:

torch.Size([3])
torch.Size([3, 1])
torch.Size([1, 3])

unsqueeze(解压)在指定位置插入一个长度为 1 的维度:

python
import torch

b = torch.zeros(3)
print(b.unsqueeze(0).shape)    # 形状 (1, 3):一行
print(b.unsqueeze(1).shape)    # 形状 (3, 1):一列

输出:

torch.Size([1, 3])
torch.Size([3, 1])

这是第 3 章遗留问题的答案:m[:, 1] 得到的 (3,) 想要保持二维,就写 m[:, 1].unsqueeze(1) 得到 (3, 1)

5.5 拼接 cat 与堆叠 stack

cat(拼接)沿已有维度把多个张量接起来,除了该维度,其他形状必须一致:

python
import torch

m1 = torch.arange(6).reshape(2, 3)
m2 = torch.arange(6, 12).reshape(2, 3)
print(torch.cat([m1, m2], dim=0))    # 纵向拼接:行数相加
print(torch.cat([m1, m2], dim=1))    # 横向拼接:列数相加

输出:

tensor([[ 0,  1,  2],
        [ 3,  4,  5],
        [ 6,  7,  8],
        [ 9, 10, 11]])
tensor([[ 0,  1,  2,  6,  7,  8],
        [ 3,  4,  5,  9, 10, 11]])

stack(堆叠)沿新插入的维度把张量摞起来,要求参与堆叠的张量形状完全相同:

python
import torch

v1 = torch.tensor([1, 2, 3])
v2 = torch.tensor([4, 5, 6])
print(torch.stack([v1, v2], dim=0))    # 新维在最前 → 2 行 3 列
print(torch.stack([v1, v2], dim=1))    # 新维在中间 → 3 行 2 列

输出:

tensor([[1, 2, 3],
        [4, 5, 6]])
tensor([[1, 4],
        [2, 5],
        [3, 6]])

记忆:cat 是「把一段接到另一段后面」,stack 是「一摞卡片叠起来,多出一个维度」。训练时把一批样本叠成 batch,用的是 stackcat

5.6 拆分 split / chunk 与展平 flatten

split(每段长度) 按固定长度切,最后一段不足就剩多少是多少;chunk(段数) 平均切成若干段:

python
import torch

c = torch.arange(10)
print(torch.chunk(c, 2))
print(torch.split(c, 3))

输出:

(tensor([0, 1, 2, 3, 4]), tensor([5, 6, 7, 8, 9]))
(tensor([0, 1, 2]), tensor([3, 4, 5]), tensor([6, 7, 8]), tensor([9]))

flatten(展平)把指定维度之后的所有维度合并成一个:

python
import torch

t = torch.arange(24).reshape(2, 3, 4)
print(t.flatten().shape)      # 全部展平 → (24,)
print(t.flatten(1).shape)     # 从第 1 维开始展平 → (2, 12)

输出:

torch.Size([24])
torch.Size([2, 12])

卷积网络接全连接层时,flatten(1) 是最常用的一步(第 10 章)。

动手实践

  1. torch.arange(12) reshape 成 (3, 4)(4, 3)(2, 2, 3),打印各自的 shape
  2. 创建 2×3 张量,转置后打印 is_contiguous();对转置结果分别调用 view(-1)reshape(-1),观察哪个报错。
  3. unsqueeze 把长度 5 的张量变成 (1, 5)(5, 1),再用 squeeze 变回来,验证形状。
  4. cat 沿 dim=0 拼接两个 3×4 张量,再沿 dim=1 拼接,分别打印形状,说明各维怎么变化。
  5. 把 4 个长度 3 的向量用 stack 分别沿 dim=0、dim=1 堆叠,比较两个结果的含义。

常见错误

错误写法现象原因
a.reshape(2, 3) 元素总数对不上RuntimeError: shape '[2, 3]' is invalid for input of size 7新形状的元素总数必须等于原总数,用 -1 让程序推断
转置后直接 view(-1)RuntimeError: view size is not compatible ... Use .reshape(...) instead.转置结果不连续,view 要求连续内存,用 reshape
torch.cat([a, b], dim=1) 时行数不同RuntimeError: Sizes of tensors must match except in dimension 1cat 除拼接维外形状必须一致
想把两个向量并成两行却用 cat变成一维长向量向量拼接用 stack(新增维度)或先 unsqueezecat
期望 m[:, 1] 是列向量形状 (3,)索引会降维,用 unsqueeze(1) 补回维度
permute 参数写错形状不对permute 参数是新维度顺序,不是目标形状

章末练习

基础

  1. 创建 torch.arange(8),用 reshape 变成 2×4,再展平回一维,验证元素顺序没变。
  2. 创建 3×4 张量并转置,打印结果形状;用 permute 对形状 (2, 3, 4) 的张量做 (1, 2, 0) 排列,打印新形状。
  3. squeeze / unsqueeze 完成:(1, 3, 1)(3,)(3, 1)(1, 3, 1),每步打印形状。

提高

  1. 解释 viewreshape 的区别,并用转置示例证明 view 在不连续张量上会报错。
  2. 把 3 个形状 (2, 4) 的张量分别用 stack(dim=0 与 dim=1)和 cat(dim=0)拼接,打印三种结果的形状并解释含义。

挑战

  1. 一个形状 (3, 4, 5) 的张量,如何得到形状 (15, 4)?列出两种可行写法并验证。
  2. split 把长度 100 的向量切成 16 段(最后一段可以不足),打印每段长度,再用 cat 拼回去验证能还原。

章末自测

每题选择一个最佳答案。本书不附答案:完成后交由老师或 AI 老师批改讲解。

  1. torch.arange(6).reshape(3, -1) 的结果形状是?
    • A. torch.Size([3, 3])
    • B. torch.Size([3, 2])
    • C. torch.Size([2, 3])
    • D. 报错
  2. 关于 viewreshape,正确的说法是?
    • A. 两者完全相同
    • B. view 要求张量连续,reshape 必要时自动复制
    • C. reshape 要求张量连续
    • D. view 总是复制数据
  3. torch.arange(6).reshape(2, 3).T.shape 是?
    • A. torch.Size([2, 3])
    • B. torch.Size([3, 2])
    • C. torch.Size([6])
    • D. torch.Size([1, 6])
  4. 对不连续张量调用 view(-1) 会?
    • A. 正常返回展平结果
    • B. 报 RuntimeError,建议改用 reshape
    • C. 静默复制数据
    • D. 报 TypeError
  5. torch.zeros(3).unsqueeze(1).shape 是?
    • A. torch.Size([3])
    • B. torch.Size([1, 3])
    • C. torch.Size([3, 1])
    • D. torch.Size([1])
  6. torch.zeros(1, 4, 1).squeeze().shape 是?
    • A. torch.Size([4])
    • B. torch.Size([1, 4, 1])
    • C. torch.Size([4, 1])
    • D. torch.Size([1, 4])
  7. 两个形状 (2, 3) 的张量做 torch.cat(..., dim=0),结果形状是?
    • A. torch.Size([2, 3])
    • B. torch.Size([4, 3])
    • C. torch.Size([2, 6])
    • D. torch.Size([6, 2])
  8. 两个形状 (2, 3) 的张量做 torch.stack(..., dim=0),结果形状是?
    • A. torch.Size([2, 3])
    • B. torch.Size([2, 2, 3])
    • C. torch.Size([4, 3])
    • D. torch.Size([2, 6])
  9. torch.arange(24).reshape(2, 3, 4).flatten(1).shape 是?
    • A. torch.Size([24])
    • B. torch.Size([2, 12])
    • C. torch.Size([6, 4])
    • D. torch.Size([8, 3])
  10. torch.chunk(torch.arange(10), 2) 的结果是?
    • A. 两个长度 5 的张量
    • B. 一个长度 10 的张量
    • C. 三个长度不等的张量
    • D. 报错