首页 / PyTorch 入门教程 / 张量操作:索引、切片、变形与广播

PyTorch 入门教程

张量操作:索引、切片、变形与广播

本教程共 60 篇 · 第 5 篇 · 更新于 2026-08-17 · 约 3 分钟阅读

张量索引切片viewreshape广播PyTorch

本节目标:学会从张量里取数、改变张量形状、把多个张量拼在一起,并理解广播(broadcasting)的规则。这些操作是后面写神经网络时天天要用的基本功。

上一章我们学会了创建张量。但光会创建没用,得会操作。训练神经网络时,你经常要从一批数据里切出某几行、把二维向量摊平成一维、把两个张量拼起来。这一章就把这些基本功一次讲清楚。

索引与切片:和列表几乎一样

如果你用过 Python 列表,张量的索引切片基本不用学。语法一样,只是多了逗号,可以同时切多个维度。

import torch

x = torch.arange(12).reshape(3, 4)
print(x)
# tensor([[ 0,  1,  2,  3],
#         [ 4,  5,  6,  7],
#         [ 8,  9, 10, 11]])

print(x[0])       # 第一行:tensor([0, 1, 2, 3])
print(x[1, 2])    # 第 2 行第 3 列:tensor(6)
print(x[:, 1])    # 所有行的第 2 列:tensor([1, 5, 9])
print(x[1:, :2])  # 后两行的前两列
print(x[-1])      # 最后一行,负索引同样适用

逗号前面管行,后面管列。冒号表示”全都要”。规则就这一条。

张量还支持用布尔掩码(mask)挑元素。先写条件,再用结果当索引,一步到位:

mask = x > 5
print(mask)
# tensor([[False, False, False, False],
#         [False, False,  True,  True],
#         [ True,  True,  True,  True]])

print(x[mask])  # tensor([ 6,  7,  8,  9, 10, 11])
Note

布尔掩码挑出来的结果是一维张量,原来的形状不会保留。数据清洗里筛掉异常值,用的就是它。

变形:view、reshape 和 -1

变形就是改形状,不改数据。最常用的是 view()reshape(),新手阶段可以当成一回事。

x = torch.arange(12)
a = x.view(3, 4)     # 改成 3 行 4 列
b = x.reshape(4, 3)  # 改成 4 行 3 列
c = x.reshape(-1)    # 展平成一维
d = x.view(2, -1)    # 2 行,列数自动推断

-1 是个偷懒神器:告诉 PyTorch”这个维度你帮我算”。只要元素总数对得上,它就能推出唯一答案。

viewreshape 有啥区别?一句话:view 要求张量在内存里连续存储,reshape 不要求,不行时它会自动拷贝一份。日常用 reshape 最省心。

还有两个常用的增删维度操作:

x = torch.randn(3, 4)
y = x.unsqueeze(0)   # 在第 0 维前加一维:torch.Size([1, 3, 4])
z = y.squeeze(0)     # 去掉第 0 维(大小为 1 才能去):torch.Size([3, 4])

图像数据经常是「1, 通道, 高, 宽」这种带批量维的格式,unsqueeze 就是给单张图加批量维的常用手段。

Tip

你也可以用 x[None, ...] 代替 x.unsqueeze(0),写法更短,两者等价。

拼接与堆叠:cat 和 stack

把多个张量合在一起,有两个容易混的 API:torch.cattorch.stack

cat 是沿着已有的某个维度拼接,张量总维度不变:

a = torch.tensor([[1, 2], [3, 4]])
b = torch.tensor([[5, 6], [7, 8]])

print(torch.cat([a, b], dim=0))  # 行方向拼:形状 (4, 2)
print(torch.cat([a, b], dim=1))  # 列方向拼:形状 (2, 4)

stack 则是新造一个维度,把张量摞进去。它要求所有张量形状完全一致:

print(torch.stack([a, b], dim=0).shape)  # torch.Size([2, 2, 2])

stack([a, b]) 可以理解成 cat([a[None], b[None]]),即先各加一维再拼接。把多个特征矩阵堆成三维张量,用的就是它。

广播:形状不同也能运算

广播(broadcasting)是 PyTorch 最方便也最容易踩坑的机制。它允许形状不同的两个张量直接做加减乘除,自动把小的”复制”成大的。

规则只有两条,比较时从右往左看:

  1. 两个维度大小相同,或其中一个是 1,就能对齐;
  2. 维度不够的一方,左边自动补 1。

看个例子:

a = torch.ones(3, 4)
b = torch.tensor([1.0, 2.0, 3.0, 4.0])  # 形状 (4,)
c = a + b
print(c)
# tensor([[2., 3., 4., 5.],
#         [2., 3., 4., 5.],
#         [2., 3., 4., 5.]])

b 形状是 (4,),左边补 1 变成 (1, 4),再复制成 (3, 4),然后逐元素相加。标量加张量也是一个道理:x + 1 就是把 1 广播到每个元素。

什么时候会报错?两个维度都不相等、也都不为 1 的时候:

a = torch.ones(3, 2)
b = torch.ones(2, 3)
# c = a + b  # RuntimeError: size mismatch
Warning

广播能少写很多代码,但也会掩盖形状错误。调试时如果数值不对劲,先打印 tensor.shape,看看是不是广播成了你没想到的形状。

形状推算:meta 设备小技巧

搭网络时经常要手算每层输出的形状。不想算的话,可以借 meta 设备跑一遍前向——不分配真实内存,只算形状:

x = torch.rand(2, 3, 10, 10, device="meta")
conv = torch.nn.Conv2d(3, 5, 2, device="meta")
out = conv(x)
print(out.shape)  # torch.Size([2, 5, 9, 9])

这个技巧在写卷积网络时特别实用,我们后面搭 CNN 时会再见到它。

这一章的四个操作——索引切片、变形、拼接、广播,就是张量的”基本功四件套”。多动手敲几遍,尤其是广播规则,最好自己造几个形状对比着试,比背规则管用。下一章进入 PyTorch 最核心的机制:自动微分(autograd)。