学习时长差不多一年、把 PyTorch 用得比较顺之后,我回过头看自己以前写的代码,发现花掉最多时间的地方不是模型设计,反而是张量操作里那些看起来简单的切分、堆叠、索引。多花的时间基本都是因为对维度没有建立起直觉:写的时候觉得没问题,一跑就报 size mismatch,再一看文档、试一波参数,小半个下午就没了。这篇文章就是我折腾完之后的一份沉淀笔记,把 PyTorch 中切分、堆叠、索引相关的 API 从参数含义、维度规则到实际选型一次讲明白。适合刚入门 PyTorch 的同学,也适合用了几个月但总被维度搞晕、看到 cat 和 stack 就犹豫的朋友。
1. 动手之前,先建立对 Tensor 维度的直觉
很多报错其实不是 API 本身复杂,而是我们根本没搞清楚"维度"这个东西在张量里到底是什么含义。切分、堆叠、索引,本质上都是在对维度做操作。所以先花点时间把维度讲透,后面全部顺理成章。
1.1 shape 就是一张多维数组的刻度表
PyTorch 里的张量(Tensor)可以理解成一个多维数组。二维张量就是我们熟悉的矩阵,有行有列;三维张量就是多个矩阵叠在一起;四维张量可以想象成多个三维立方体组成的序列。
shape 属性就是这张多维数组的刻度表。torch.Size([3, 64, 64]) 这个形状从左到右依次代表第 0 维、第 1 维、第 2 维的大小。其中第 0 维是 3,第 1 维是 64,第 2 维是 64。
在 PyTorch 的视觉任务里,[3, 64, 64] 通常表示一张 3 通道、高 64 像素、宽 64 像素的图片。注意通道数排在最前面,这跟 OpenCV、PIL 经常看到的 HWC 排布(高、宽、通道)是不一样的。PyTorch 默认使用 [C, H, W],批量数据就是 [B, C, H, W]。记住这个排布习惯,后面很多操作我们都可以根据它来推断轴的方向和含义。
python复制import torch
x = torch.randn(3, 64, 64)
print(x.shape)
# torch.Size([3, 64, 64])
1.2 axis 编号从 0 开始,和 shape 位置一一对应
PyTorch 的绝大多数 API 都接受一个 dim 参数,表示你要操作哪一条轴。对二维矩阵来说,dim=0 是行方向,dim=1 是列方向。对四维卷积数据 [B, C, H, W] 来说,dim=0 是 batch 方向,dim=1 是通道方向,dim=2 是高方向,dim=3 是宽方向。
举个例子,假设有两个形状都是 [2, 3] 的张量:
- 沿
dim=0拼接,相当于把矩阵上下摞起来,新形状是[4, 3]; - 沿
dim=1拼接,相当于把矩阵左右并起来,新形状是[2, 6]。
这个例子很好地说明了"沿着哪条轴操作,哪条轴的大小会变化"。后面的切分和堆叠,绝大多数都可以用这个规则去理解和推导。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 切分操作:chunk、split、unbind 的选型思路
切分就是把一个张量沿着某个维度拆成多个部分。PyTorch 给了我们三种常用切分 API:torch.chunk、torch.split、torch.unbind。它们看着像,但适用场景差别很大。
2.1 torch.chunk:按份数均分,但没那么"均"
torch.chunk(input, chunks, dim=0) 接收一个份数,语义上表示"我要把这个张量沿某条轴分成几份"。
代码看起来很简单:
python复制import torch
x = torch.arange(10)
c1, c2 = torch.chunk(x, 2)
print(c1)
# tensor([0, 1, 2, 3, 4])
print(c2)
# tensor([5, 6, 7, 8, 9])
但有个细节坑过不少新手:当总长度不能被份数整除时,chunk 并不会报错,而是会把前面的块切得稍大一点,最后一块可能比前面小。例如 torch.chunk(torch.arange(10), 3),因为 10 / 3 = 3.33,向上取整得到 4,所以第一块和第二块都会分到 4 个元素,最后一块只有 2 个元素。
这个行为意味着,如果你依赖"每块大小完全一样"来做后续操作,chunk 不是最合适的选择。比如做模型并行切分、或者将一批数据均匀分给多个 worker,建议先验证整除性,或者直接用 split。
2.2 torch.split:按长度切分,可控性更强
torch.split(tensor, split_size_or_sections, dim=0) 的核心参数不是份数,而是"每一块的长度"。
它有两种用法:
第一种,传入一个整数,表示每一块大小:
python复制x = torch.arange(10)
s1, s2, s3 = torch.split(x, 4)
print(s1.shape, s2.shape, s3.shape)
# torch.Size([4]) torch.Size([4]) torch.Size([2])
可以看到,按长度 4 去切,10 个元素会被切成 4、4、2,最后剩下的不足部分单独成一块。
第二种,传入一个列表,精确指定每一块的长度:
python复制x = torch.arange(10)
s1, s2, s3 = torch.split(x, [3, 3, 4])
print(s1, s2, s3)
# tensor([0, 1, 2]) tensor([3, 4, 5]) tensor([6, 7, 8, 9])
使用列表时,列表内的所有数字加起来必须等于该维度的大小,否则会报错。
所以在需要"每块大小固定值"、或者"每块大小不一样"的场景,split 明显比 chunk 更好用。
2.3 torch.unbind:把一个维度整个拆掉
torch.unbind(tensor, dim=0) 的作用是删除指定维度,并把该维度上的每个切片作为独立张量返回,返回结果是一个元组。
python复制x = torch.randn(3, 4)
a, b, c = torch.unbind(x, dim=0)
print(a.shape)
# torch.Size([4])
这个操作相当于 split(x, 1, dim=0) 的快捷方式,但它更直接:直接去掉这一维,得到的是降维后的张量。
它在处理 batch 上的用途很广。训练好的模型做推理时,如果你想逐个样本看输出结果,用 unbind 就比 chunk(x, batch_size) 然后再取下标干净得多。
2.4 三个 API 的对比和选择逻辑
| API | 参数核心 | 返回结果 | 最佳场景 |
|---|---|---|---|
torch.chunk |
分成几份 | 每份是原张量的视图,块数不超过指定份数 | 只关心份数、不关心具体大小的场景 |
torch.split |
每块长度或长度列表 | 每份是视图,长度严格受控 | 需要精确控制每块长度时 |
torch.unbind |
按维度拆开 | 返回元组,且维度降低 | 要按某轴逐条取样本 |
我的习惯是:只要涉及精确分配长度,一律用 split;只在极少数明确要求"分成几份"的场景用 chunk;而"把 batch 里的每个样本单独拿出来"这个动作,直接用 unbind,比前面两者都更符合直觉。
3. 堆叠操作:cat、stack、concatenate 的维度逻辑
切分和堆叠就像一体两面。切分是把一条轴拆小,堆叠则是把多个张量组合起来。PyTorch 里最常见的是 torch.cat 和 torch.stack,另外还有不少人会遇到 torch.concatenate。
3.1 cat:沿已有维度拼接,不新增维度
torch.cat(tensors, dim=0) 的意思是把输入张量列表沿着某个已有维度拼接起来。拼接后,除了被拼接的这条轴,其他维度的大小必须完全一样。
python复制a = torch.randn(2, 3)
b = torch.randn(4, 3)
c = torch.cat([a, b], dim=0)
print(c.shape)
# torch.Size([6, 3])
如果两个张量在被拼接维度之外的大小不一致,会直接报错。比如一个 [2, 3] 和一个 [2, 4] 沿 dim=0 拼接,因为第 1 维分别是 3 和 4,不匹配,所以会失败。
cat 适用于特征拼接、按通道拼接图片等场景。比如你有两个形状为 [4, 16, 8, 8] 的特征图,想把它们按通道维度拼起来得到 [4, 32, 8, 8],用 torch.cat(..., dim=1) 就行。
python复制f1 = torch.randn(4, 16, 8, 8)
f2 = torch.randn(4, 16, 8, 8)
out = torch.cat([f1, f2], dim=1)
print(out.shape)
# torch.Size([4, 32, 8, 8])
3.2 stack:沿新增维度堆叠,会插入一个新轴
torch.stack(tensors, dim=0) 最大的特点是:它不会沿着已有维度拼接,而是在指定位置插入一个新维度。因此它要求所有输入张量的形状完全一致,因为新维度的大小等于你传入张量的数量。
python复制a = torch.randn(3, 4)
b = torch.randn(3, 4)
s0 = torch.stack([a, b], dim=0)
s1 = torch.stack([a, b], dim=1)
s2 = torch.stack([a, b], dim=2)
print(s0.shape) # torch.Size([2, 3, 4])
print(s1.shape) # torch.Size([3, 2, 4])
print(s2.shape) # torch.Size([3, 4, 2])
注意观察:同样是两个 [3, 4] 的张量,stack 后新维度的大小恒为 2(张量数量),同时原来所有维度仍然保留。而 cat 则不会增加维度数量,只是让某一维变大。
这个"插入新维"的特性,让 stack 在构造 batch 时非常常用。比如某个数据集中每张图片的特征向量是 [10] 的形状,现在要把 5 个图片特征组装成一个 batch,就需要 torch.stack(..., dim=0),得到 [5, 10]。如果这里误用 cat,得到的是 [50],batch 维度被拍平,模型根本没法处理。
3.3 torch.concatenate 和 torch.cat 是什么关系
torch.concatenate 几乎是 torch.cat 的同一回事,它在 PyTorch 后续版本中作为更贴近 NumPy 命名习惯的别名形式提供。它的功能、参数和返回值跟 cat 没有区别:
python复制c1 = torch.cat([a, b], dim=0)
c2 = torch.concatenate([a, b], dim=0)
# c1 和 c2 的结果完全一致
所以你在代码里看到这两种写法都不用慌,本质上是同一个操作。个人建议新代码统一使用 torch.cat,因为社区里绝大多数代码和文档都使用这个名字,可读性更高。
3.4 什么时候用 cat,什么时候用 stack
这里有一个很实用的判断标准:
- 如果你要组合的张量已经存在那条"容纳它们的轴",就用
cat; - 如果你希望把它们"整体打包",新增一个维度来装它们,就用
stack。
打个比方,cat 像是把几本同规格的书排列在同一个书架上,书架层数不增加;stack 像是每本书放进一个独立格子,再把这些格子摞起来,整个结构多了一层。
| 操作 | 是否新增维度 | 对输入形状要求 | 典型场景 |
|---|---|---|---|
cat |
否 | 仅拼接维度可变,其他维度必须一致 | 特征拼接、通道拼接 |
stack |
是 | 所有维度必须完全一致 | 构造 batch、把多个样本堆出一个新维度 |
4. 索引体系:切片、布尔掩码和花式索引
切分、堆叠之外,索引是处理张量数据时最常用的"隐形切分"手段。它不止是 Python 列表索引的简单延伸,还多了些张量独有的规则。
4.1 基础切片规则与负数索引
PyTorch 多维张量的切片遵循"每个维度用逗号隔开"的规则。x[0:2, 1:3] 表示第 0 维取 0 到 1 行,第 1 维取 1 到 2 列。
python复制x = torch.arange(12).reshape(3, 4)
print(x)
# tensor([[ 0, 1, 2, 3],
# [ 4, 5, 6, 7],
# [ 8, 9, 10, 11]])
print(x[1:, 1:3])
# tensor([[ 5, 6],
# [ 9, 10]])
负数索引从末尾开始计数。x[-1] 取最后一行,x[:, -2:] 取每一行的最后两列。
对于索引结果,有一个新手非常容易忽略的差异:
x[:, 0]会降维,形状从[3, 4]变成[3];x[:, 0:1]会保持维度,形状从[3, 4]变成[3, 1]。
这两个结果看起来都是"取第一列",但维度数量不同。这个问题在送入 RNN、Transformer 等对维度敏感的网络时经常暴露出来。我的建议是:书写索引前先想清楚"我到底是要压掉一维,还是保留这一维结构",再决定用整数索引还是切片索引。
4.2 布尔掩码索引:按条件筛选数据
布尔掩码是索引里最实用的功能。它允许你用一个形状相同的布尔张量做筛选,True 的位置对应的元素会被取出来。
python复制x = torch.tensor([-1, 2, -3, 4])
mask = x > 0
print(mask)
# tensor([False, True, False, True])
print(x[mask])
# tensor([2, 4])
更高维的用法也很干净。比如你要把标签等于 0 的样本全部筛出来:
python复制x = torch.randn(6, 3, 32, 32)
y = torch.tensor([0, 1, 0, 2, 1, 0])
mask = y == 0
x_selected = x[mask]
print(x_selected.shape)
# torch.Size([3, 3, 32, 32])
布尔掩码返回的结果通常是一维张量,或者不规则的形状,而且返回的是数据的拷贝,不是视图。因此后续对它修改不会影响原张量,这在数据预处理中通常是我们想要的行为。
4.3 整数数组索引(花式索引):按索引列表取值
整数数组索引,也叫"花式索引",允许你传入一个索引数组,一次性取出多个不连续位置上的元素。
python复制x = torch.arange(12).reshape(3, 4)
row_idx = torch.tensor([0, 2])
print(x[row_idx])
# tensor([[ 0, 1, 2, 3],
# [ 8, 9, 10, 11]])
也可以同时索引多个维度:
python复制print(x[torch.tensor([0, 1]), torch.tensor([1, 2])])
# tensor([1, 6])
# 取第0行第1列、第1行第2列
花式索引也是拷贝,不是视图。这在处理数据打乱、按采样索引取子集时很有用。
4.4 索引和切分堆叠的联动
索引本身就可以替代很多 chunk 操作。比如你想从一个大 batch 中按某个自定义规则挑出部分样本,再重新组成一个小 batch,用索引加 stack 就非常顺畅:
python复制x = torch.randn(10, 3, 32, 32)
idx = torch.tensor([3, 7, 9])
subset = x[idx]
print(subset.shape)
# torch.Size([3, 3, 32, 32])
这比先用 chunk 拆出 10 份、再手动挑选要简洁得多。索引配合 cat、stack,可以覆盖绝大部分数据筛选、重组的场景。
5. 把这些 API 串成一个数据预处理流程
在单个 API 层面理解之后,最好把它们放进一个更接近实际任务的小流程里跑一遍。这里我设计一个非常典型的需求:把一个 batch 里的图片,按类别拆开,分别做不同的归一化,再重新合并回一个 batch。步骤不算复杂,但切分、掩码索引、堆叠都会用上。
5.1 场景设定
假设有 8 张图片,形状为 [8, 3, 32, 32],对应的标签是长度为 8 的张量。我们希望把标签为 0 的图片归一化到 [-1, 1],标签为 1 的图片归一化到 [0, 1],然后再合并起来。
直接使用 unbind 按图片拆开,逐张判断处理,最后 stack 回去,是一种比较直观的写法。但这种写法在 batch 较大时效率不高。更高效的方式是用布尔掩码直接筛选出两个子集,分别处理,再 cat 回去。
5.2 实现步骤
python复制import torch
x = torch.randn(8, 3, 32, 32)
y = torch.tensor([0, 1, 0, 0, 1, 1, 0, 1])
mask_0 = y == 0
mask_1 = y == 1
x0 = x[mask_0]
x1 = x[mask_1]
# 标签为0的样本归一化到 [0, 1]
x0 = (x0 - x0.min()) / (x0.max() - x0.min())
# 标签为1的样本归一化到 [-1, 1]
x1 = 2 * x1 - 1
# 重新合并成一个 batch
x_combined = torch.cat([x0, x1], dim=0)
print(x_combined.shape)
# torch.Size([8, 3, 32, 32])
这里有几个地方需要提醒:
x[mask_0]后的x0第一个维度大小是动态的,可能不是 4,取决于实际标签分布;cat要求除了被拼接的维度,其余维度完全一致。这里两个子集都是 3 通道 32×32 图片,所以可以安全拼接;- 拼接后顺序发生变化。合并结果里前面的样本全是标签 0,后面全是标签 1。如果后续要使用训练循环里的 mini-batch 随机采样,需要额外做一次打乱,否则每个 batch 的类别分布会非常不均衡。
5.3 维度变化跟踪
我们跟踪一下这个流程里的形状变化:
- 原 batch:
[8, 3, 32, 32] - 布尔掩码索引后:
x0变为[N0, 3, 32, 32],x1变为[N1, 3, 32, 32],其中N0 + N1 = 8 - 经过不同归一化处理,形状不变
cat拼接后:[8, 3, 32, 32]
这种"先筛选、后合并"的模式,在数据增广、类别均衡采样、多任务分支处理里非常常见。尤其是对不同类别使用不同预处理逻辑时,布尔掩码加 cat 的组合比用 for 循环逐张处理要高效,代码也更简洁。
6. 实操中踩过的坑与排查经验
最后这部分是我在大量数据预处理、断点调试过程中积累下来的真实问题。有些问题不致命,但足以消耗大量时间;有些问题比较隐蔽,甚至会悄悄污染数据。
6.1 视图与副本:共享内存是双刃剑
普通切片(如 x[0:2])是视图,返回的张量与原张量共享内存;布尔掩码索引和花式索引则是副本,不共享内存。
这个区别放在实际代码里非常容易踩雷。举个例子,你写了这样一段:
python复制x = torch.arange(10)
sub = x[0:3]
sub[0] = 99
print(x)
# tensor([99, 1, 2, 3, 4, 5, 6, 7, 8, 9])
因为 sub 是视图,改 sub 会把原数组也改了。如果代码里后续还有别的统计逻辑,这个"悄悄"的改动会导致诡异的问题,而且非常难定位。
反过来,如果每次都用 x[torch.tensor([0, 1, 2])] 这种花式索引,又会频繁复制数据,内存开销变大。我的习惯是:只想读不改,两种都行;要改原数据,用视图切片;想得到一份独立数据再随便折腾,用 clone() 或者高级索引。
python复制# 如果确实要一个独立副本
sub = x[0:3].clone()
sub[0] = 99
print(x)
# tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9])
6.2 设备不一致导致的异常操作
切分、堆叠、索引本身不涉及计算,但组合使用时,很容易出现"CPU 张量和 GPU 张量混在一起"的问题。
比如从 DataLoader 取出来的 batch 在 GPU 上,你用掩码索引筛选子集后得到的还是 GPU 张量。如果你紧接着把它和一个 CPU 张量做 torch.stack,或者做某些元素级运算,PyTorch 会直接报错。
排查手段很简单:打印 .device 属性。
python复制print(x0.device)
# cuda:0 或 cpu
拿到 GPU 张量后,想和其他操作统一设备,建议通过 .to(device) 做显式转换,而不是依赖隐式转换,否则在训练循环里很容易出现某些步骤移动了张量、某些步骤没移动,代码运行到一半才报错的情况。
6.3 动态形状导致 stack 失败
布尔掩码索引的结果形状是动态的,因为它依赖实际满足条件的样本数量。如果想对两个类别的数据用 stack 堆叠成新 batch,一旦两个子集的数量不同,stack 会因为形状不一致直接报错。
此时有两个方向:
- 用
cat而不是stack,因为cat只要求拼接维度之外完全一致; - 如果确实需要每个样本有独立的 batch 维度,可以先把数据
pad到相同的补齐长度,再做stack。
实际写代码时,我一般会先打量一下后续模型输入的要求。模型通常只要求一个 batch 维,不强制要求每个"来源"单独成一维。所以用 cat 拼接是默认选择,只有明确需要"保留批次来源"时才用 stack。
6.4 索引后梯度与原地修改的问题
切分、堆叠、索引这些操作基本都是可导的,在自动求导框架里可以放心用在网络内部。但要注意:如果对这些张量做原地修改(比如 x0 += 1),可能会影响反向传播中的梯度计算。
尤其是当这个张量是从某个需要梯度的中间结果切片来的,原地修改可能导致梯度不再准确。PyTorch 的 autograd 在检测到原地操作破坏计算图时,有时会直接报 RuntimeError,有时则不会。最稳妥的方式是避免在需要梯度的中间张量上直接原地修改,改成构造新张量:
python复制# 不推荐
x0 += 1
# 推荐
x0 = x0 + 1
这两种写法数值结果一样,但后者不会破坏计算图,潜藏的坑少得多。
6.5 一个小习惯:先手推 shape 再写代码
最后分享一个我自己的小习惯。写任何包含切分、堆叠、索引的数据预处理逻辑前,我会先在草稿纸上写下输入形状,然后一步步手推每个中间结果应该是什么形状,最后再去写代码。别小看这一步,它帮我省下的调试时间非常可观。很多 size mismatch 的报错,其实在动手写代码之前就已经能被预测到。
如果对不确定的 API 行为有疑问,可以快速甩几个随机张量测试一下:
python复制a = torch.randn(4, 3)
b = torch.randn(4, 3)
print(torch.stack([a, b], dim=1).shape)
# torch.Size([4, 2, 3])
这个方法比反复翻文档快得多,也更能理解 API 的维度逻辑。等你熟练掌握了这套"手推 + 快速验证"的组合,再面对高维张量的切分、堆叠、索引时,基本不会再被绕晕。
