张量操作:索引、切片、变形与广播
本教程共 60 篇 · 第 5 篇 · 更新于 2026-08-17 · 约 3 分钟阅读
本节目标:学会从张量里取数、改变张量形状、把多个张量拼在一起,并理解广播(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”这个维度你帮我算”。只要元素总数对得上,它就能推出唯一答案。
那 view 和 reshape 有啥区别?一句话: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.cat 和 torch.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。
看个例子:
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)。