开头先说一下为什么要把这几个操作单独拿出来写一篇。不管是做CV还是做NLP,也不管是在搭网络还是在写数据加载逻辑,PyTorch里最绕不开的就是张量操作。我见过不少人背了一遍API就开始写模型,结果一遇到形状对不上、维度爆炸、广播规则搞不清就开始懵。这篇笔记就围绕索引分片、合并和维度调整这三块,把常用的方法原理、实际用法和你以后一定会踩的坑讲清楚。内容按“先定位元素—再组合数据—最后改形状”的顺序推进,新手可以顺着学,有基础的朋友直接跳到自己薄弱的部分查漏补缺。
1. 索引与分片:先掌握张量的“定位”能力
1.1 从Python列表说起:NumPy风格还是PyTorch风格?
PyTorch的张量索引整体延续了NumPy的语法风格。定义一个二维张量:
import torch x = torch.arange(20).reshape(4, 5) # tensor([[ 0, 1, 2, 3, 4], # [ 5, 6, 7, 8, 9], # [10, 11, 12, 13, 14], # [15, 16, 17, 18, 19]])最基本的整数索引和Python列表类似,x[0]拿到第一行,x[0, 1]拿到第一行第二列的元素(从0开始计数)。但这里有一个非常容易忽略的区别:x[0]返回值是张量而不是列表视图,它的维度是(5,)。如果你写x[0][1],虽然也能取到元素,但写代码时建议把索引放在同一对方括号里:x[0, 1]。两种写法结果一样,但后者更符合张量的语义,切分的表达力也更强。
Python列表的切片是list[start:stop:step],PyTorch张量的切片语法同理,只是可以针对每个维度各写一段切片。比如:
x[1:, 2:] # 从第1行到最后一行、第2列到最后一列 x[:, ::2] # 所有行、从第0列开始每隔一列取一列这里要注意一个和Python列表的差异:切片结果和原张量共享内存。也就是说,对切片结果做原地修改,会影响原始数据。这一点既像NumPy,又比Python列表更“危险”。很多人写数据增强时,在切片上做归一化,以为只是改了副本,结果原始数据也被改了,再拿来算损失就全错了。除了切片,后面讲的view、transpose等操作也有这类共享内存问题,建议养成一个习惯:只要打算后续修改数据,就显式调一次.clone()。
1.2 花式索引(Fancy Indexing)的原理与踩坑点
花式索引是用张量/列表作为索引去取数据。它可以按行取、按列取,也可以同时按行列取。直接看代码:
indices = torch.tensor([0, 2, 3]) x[indices] # 取第0、2、3行 x[:, indices] # 取所有行的第0、2、3列 x[indices, indices] # 取(0,0), (2,2), (3,3)三个位置的元素这里的重点在最后一行:x[indices, indices]并不是取“行索引和列索引的笛卡尔积”,而是把两个索引张量按位置一一配对。想要笛卡尔积,得用组合索引或者torch.meshgrid先构建对应的矩阵。
花式索引的结果和切片不一样,它会把数据复制一份,得到新张量,不和原张量共享内存。这点刚入门很容易搞混。切片是“视图”,修改会影响原数据;花式索引是“拷贝”,修改不影响原数据。
还有一个小细节,花式索引传Python列表和传张量的差别。x[[0, 2, 3]]和x[torch.tensor([0, 2, 3])]拿到的数据一样,但结果的类型和梯度传播路径是相同的。真正要注意的是,当索引值超出了张量维度范围,运行时会直接报错IndexError: index 5 is out of bounds,这一点比NumPy更严格,NumPy在部分版本下会产生警告但不一定报错。PyTorch的严格模式其实是好事,越早发现问题,越不容易留下隐藏bug。
1.3 布尔索引:实战中最好用的筛选工具
布尔索引在数据预处理阶段用得非常多。它的核心思路是:用条件和原张量形状相同的布尔张量,返回所有为True的位置上的元素。几个典型场景:
x = torch.tensor([1.0, -2.0, 3.5, -0.5]) mask = x > 0 x[mask] # 取所有正数 x[x < 0] = 0 # 把所有负数原地置为0 # 多维场景:找出所有大于10的元素 x2 = torch.randint(0, 20, (4, 5)) rows = torch.nonzero(x2 > 10) # 返回满足条件的坐标布尔索引的结果也总是进行数据拷贝。还有一个高频技巧是结合多个条件:x[(x > 0) & (x < 10)],注意括号不能省。很多人写成x[x > 0 and x < 10],直接报错,因为Python的and不能直接用在张量上,必须逐元素逻辑运算&、|。
对了,想定位坐标时不要用torch.where去“取元素”,它更适合做三目运算。要拿到满足条件位置的坐标,直接用torch.nonzero(返回形状为[非零个数, 维度数])。如果想要“第0维为行索引,第1维为列索引”的老式返回结果,就调torch.nonzero(..., as_tuple=False),这个参数默认状态在新版PyTorch里已经调整过,使用前看一眼版本说明。
2. 合并:把多个张量拼在一起
2.1 torch.cat拼接:沿已有维度把数据接起来
合并操作最常用的就是torch.cat和torch.stack两个函数,先说说torch.cat。它的作用是在已有的某个维度上把多个张量拼接在一起,好比把几段绳子头尾相连。
a = torch.randn(2, 3) b = torch.randn(4, 3) c = torch.cat([a, b], dim=0) # 形状(6, 3) # 如果做dim=1,需要两个张量第0维相同 d = torch.cat([a, a], dim=1) # 形状(2, 6)对于dim=0的拼接,除了第0维以外的所有维度都必须完全一致;对于dim=1的拼接,除了第1维以外的所有维度都必须一致。这个限制很好理解:拼接就像盖楼时把结构件对接,只有接口截面一致才能严丝合缝。如果a是(2, 3),b是(2, 4),你硬要做dim=1,运行时会报错:
RuntimeError: Sizes of tensors must match except in dimension 1. Expected size 3 but got size 4这段报错信息几乎是所有新手都会遇到的。解决办法就是先检查形状,再决定拼接的维度。多数情况下你真正需要的是“把不同样本拼成一个batch”,也就是在dim=0上拼,这时只要保持每个样本的内部维度一致即可。
在实际项目中,torch.cat最常见的用法是在DataLoader的collate_fn里处理变长数据。比如NLP里一句话七个词,另一句话九个词,没法直接拼成(2, seq_len, hidden),那就先pad到相同长度再cat。还有一个常见场景是把多尺度特征图在通道维上拼起来(U-Net的skip connection就是典型),这时就要用dim=1。
2.2 torch.stack合并:凭空多出一个维度
torch.stack和torch.cat最大的区别在于:stack不会沿着已有维度拼接,而是把所有输入张量堆叠在一个新维度上。
a = torch.zeros(2, 3) b = torch.ones(2, 3) c = torch.stack([a, b], dim=0) # 形状(2, 2, 3) d = torch.stack([a, b], dim=1) # 形状(2, 2, 3) e = torch.stack([a, b], dim=2) # 形状(2, 3, 2)以dim=0为例,新的第0维的尺寸就是输入张量的个数。既然是新维度,那么要求所有输入张量的形状完全相同。cat允许某个维度不同,stack则要求全部一致。
什么时候用stack呢?最典型的就是把一组独立维度相同的张量合成batch。比如你的数据是[特征向量],每条形状都是(10,),有32条数据,你希望的batch形状是(32, 10),用torch.stack(data_list)直接搞定。如果用torch.cat,两个(10,)在dim=0上只能得到(20,),维度就少了一层。
还有一个小技巧:想给张量批量增加一个维度,例如把(3, 4, 5)变成(3, 1, 4, 5)时也可以用torch.stack,把同一个张量放两次再选其中一个切片即可,但更干净的做法是直接unsqueeze,后面专门讲。
2.3 合并操作的常见报错与处理思路
合并时比较常见的两个报错:
第一个是“形状不一致”报错。cat要求非目标维度完全一致,stack要求所有维度完全一致。处理思路就一句话:看报错里提到的维度数字,再回去打印两个张量的.shape,逐一比对。
第二个是“数据类型不一致”报错。合并时,PyTorch要求参与操作的所有张量dtype必须一样。比如一个张量是torch.float32,另一个是torch.float64,直接报错。这时要么统一转成float32,要么都转成float64(在CPU上还会涉及内存翻倍的问题,尽量统一用float32)。
另外还要注意一个容易忽略的坑:合并操作不是原地修改,不会改变原张量。很多人写torch.cat([a, b], dim=0),然后打印a.shape,发现没变化,以为拼接失败了。其实cat返回新张量,必须用c = torch.cat([a, b], dim=0)这种形式接住返回值。这个概念在PyTorch里无处不在,很多操作都不会原地改输入,不仅限于合并。
3. 维度调整:改变张量的形状
3.1 view与reshape的关系与本质区别
维度调整是PyTorch里最让人迷惑的部分,尤其是view和reshape的区别。简单说:view要求张量在内存中满足“连续”(contiguous)条件,reshape不要求,reshape在底层会自动复制数据使其连续。
先看代码:
x = torch.arange(12) y = x.view(3, 4) # 形状(12,) -> (3, 4) z = x.reshape(3, 4) # 同样结果,但处理机制不同如果张量已经是连续的,两者没有任何区别。但当张量做过transpose、permute等操作后,内存布局会变得不连续,此时调用view会直接报错:
RuntimeError: view size is not compatible with input tensor's size and stride (...)遇到这种报错,最简单的处理办法是x.contiguous().view(...),或者直接用reshape。
这里的关键是理解“逻辑形状”和“内存布局”的区别。你可以把张量想象成一本页码连续的字典,view只是改变了你读字典的方式,但要求页码必须仍然从0到N连续排列;一旦你调了transpose把书的章节顺序打乱了,页码对不上,view就失效了,必须先把书重新整理成新的连续页码(contiguous)再操作。
我在项目里一直保持这样的习惯:只要不是对性能极其敏感的核心循环,优先用reshape而不是view。因为reshape的语义更接近“我要一个形状上的新张量”,不用操心底层连续性问题。但如果你做张量并行或者高频处理,就要注意reshape在非连续场景下会引入一次拷贝,这个开销在某些场景下不可小视。
3.2 permute与transpose:维度怎么换才不丢数据
transpose和permute都是交换维度,但transpose一次只交换两个维度,permute可以任意排列所有维度。看示例:
x = torch.randn(2, 3, 4) x.transpose(0, 2) # 交换0维和2维,形状变成(4, 3, 2) x.permute(2, 1, 0) # 把原2维放前面、原1维放中间、原0维放最后,形状同样(4, 3, 2)如果只是交换两个维度,transpose(0, 2)和permute(2, 1, 0)结果不完全相同:permute实际上可以同时完成多维交换,最终形状取决于你给出的轴顺序。记住一个口诀:permute(原维度的新顺序),传入的参数表示“新张量的第i维对应原张量的第几个维”。
从数据角度看,permute和transpose都不会改变数据的值,只是改变了读取顺序。它们返回的往往是不连续的张量,所以在后续接view前记得先contiguous()。
transpose还有一个常见误用:x.T。T在PyTorch里是2D张量的转置快捷写法,在更高维张量上,x.T只能做维度完全反转,效果等同于x.permute(*reversed(range(x.dim())))。如果你只想交换两个维度,用transpose更直观。
在图像处理里,transpose是最常见的“通道转换”工具。OpenCV读出来是H, W, C(128, 128, 3),要转成PyTorch模型需要的C, H, W(3, 128, 128),一行img.transpose(2, 0, 1)搞定。
3.3 unsqueeze、squeeze与flatten的使用场景
unsqueeze是在指定位置插入尺寸为1的新维度,squeeze是删除所有尺寸为1的维度(或指定删除某个维度)。flatten则是把连续若干维压成一维。
x = torch.randn(3, 4) x.unsqueeze(0) # 形状(1, 3, 4),类似增加batch维度 x.unsqueeze(1) # 形状(3, 1, 4) x.unsqueeze(2) # 形状(3, 4, 1) x.squeeze(0) # 如果第0维是1才会删除,否则不变 x.reshape(-1) # 展平成一维,等价于flattenunsqueeze在深度学习里有个极为经典的场景:单条样本转batch。预测时你有一张图片,模型要求输入(N, C, H, W),但当前只有(C, H, W),直接img.unsqueeze(0)就能得到(1, C, H, W)。
flatten和reshape(-1)的结果通常是一样的,但flatten可以指定起止维度。nn.Flatten(start_dim=1)就是把从第1维开始往后全部压成一维,比如(32, 3, 224, 224)变成(32, 3 * 224 * 224)。全连接层前接特征图时,nn.Flatten()是标配。
还有一个细节:squeeze不加参数时,会删除所有尺寸为1的维度。这可能会引入一个隐藏bug——如果你本意只是删掉某一个维度,但其他维度恰好也是1,结果就会比预期多删。所以需要精确控制时,一定要传维度参数:x.squeeze(0)。
4. 综合实战:维度变换在真实项目中的套路
4.1 典型流程:从加载数据到送入模型
把上面三个部分串起来看一个完整的例子。假设你在做一个人脸关键点检测任务,数据加载器返回的样本形状是这样的:
- 图像:
(H=128, W=128, C=3)(OpenCV读入,通道在最后) - 关键点:
(68, 2)(68个点的x/y坐标)
模型输入要求是图像(batch, C=3, H=128, W=128),关键点标签要求是(batch, 68*2)。
import cv2 import torch # 模拟一张图像和关键点 img = cv2.imread("face.jpg") # (128, 128, 3) pts = torch.randn(68, 2) # (68, 2) # 通道调整:HWC -> CHW img_t = torch.from_numpy(img).permute(2, 0, 1) # (3, 128, 128) # 单样本转batch img_batch = img_t.unsqueeze(0) # (1, 3, 128, 128) # 关键点展平 pts_batch = pts.reshape(1, -1) # (1, 136) # 多个样本合并成一个大batch,用torch.cat batch_imgs = torch.cat([img_batch, img_batch], dim=0) # (2, 3, 128, 128)这个流程里你同时用到了permute(换轴)、unsqueeze(加维度)、reshape(展平)和cat(合并)。写模型时会发现,几乎每个数据加载流程都会重复这几个操作。把这些套路记住,比死背API有用得多。
4.2 从PyTorch到ONNX:维度兼容问题
很多模型训练完之后要导出ONNX(热词里也有大量pytorch转onnx的搜索),这时维度调整就显得更重要了。ONNX导出时会固定输入张量的shape,但允许动态维度。如果你在模型里使用了不规范的view、reshape,导出的ONNX计算图里生成的Transpose、Reshape节点就可能会多出很多,推理框架(如OnnxRuntime)跑起来性能就差。
实践上建议:
- 模型内部尽量用
reshape代替view,但在知道张量连续时也用view保证效率; - 动态维度场景(比如可变batch size)用
torch.onnx.export时设置动态轴; - 如果模型里包含
flatten,导出的图中通常会有对应的Reshape节点,这是正常的,但要注意某些推理硬件对动态reshape支持不好,能固定shape就固定。
我踩过的一个实际坑是:模型里用了x.view(x.size(0), -1),当batch size变化时ONNX导出没问题,但在一些NPU设备上推理会报错。换成x.reshape(x.size(0), -1)后同样报错,最后改成显式x.contiguous().view(...)才稳定。不是view本身有问题,而是部分推理框架对“非连续视图”支持有限,导出前最好在模型里显式保证连续性。
4.3 广播机制的联动:维度调整和加法一起用
张量做加减法时,维度不同但满足广播规则也能操作。广播规则概括起来是:从最后一个维度依次往前比对,两个维度相同或者其中一个是1,就能对齐。例如(3, 1)和(1, 4)可以相加得到(3, 4)。
这个规则和维度调整经常配合使用。比如你想给一个batch的每个样本加一个bias向量:
data = torch.randn(32, 10) # batch, feature bias = torch.randn(10) # 一维向量 out = data + bias # 自动广播,(32, 10) + (10,) -> (32, 10)但如果bias的维度是(10, 1),就不能直接加了,需要bias.squeeze(1)或bias.reshape(10)。反过来,如果特征是二维特征图(32, 3, 224, 224),要加一个针对通道的均值(3,),就要把它变成(1, 3, 1, 1)再相加,常用写法是mean[None, :, None, None],等价于mean.unsqueeze(0).unsqueeze(2).unsqueeze(3)。
广播虽然方便,但也是最容易出隐藏bug的地方。特别是两个张量形状不匹配但广播后得到意外结果时,不会报错,只会输出一个更大形状的张量,后续计算看着没问题,实际已经错了。遇到这种情况的最快排查方式:在+、*操作后立刻打印out.shape,只花一行代码的时间,能省一整晚的定位时间。
5. 常见报错与排查心得
5.1 报错信息速查表
| 报错信息关键词 | 原因 | 处理方式 |
|---|---|---|
Sizes of tensors must match | cat拼接维度不一致 | 打印两个shape,确认非拼接维度完全一致 |
view size is not compatible | 对非连续张量用了view | 先contiguous()再view,或直接用reshape |
index 5 is out of bounds for dimension | 索引超出范围 | 检查索引张量最大值与对应维度大小 |
Expected all tensors to be on the same device | 张量跨设备 | 统一to(device) |
Expected scalar type Float but found Double | dtype不匹配 | 统一.float()或.double() |
Boolean indexing requires matching shapes | 布尔索引mask尺寸不匹配 | 检查mask与目标张量形状 |
这张表基本覆盖了索引、合并、维度调整这三类操作的大多数报错。记牢关键词之后,能少走很多弯路。
5.2 自动求导背景下的维度陷阱
在requires_grad=True的张量上做维度调整,大多数情况下梯度能正常传播,但有几个坑要注意。
第一,不是所有操作都能保持梯度。torch.nonzero和布尔索引这类操作,虽然可以用于前向计算,但涉及到的索引是不可微的。如果它出现在网络中间,梯度回传会在这一步断掉。例如:
x = torch.randn(10, requires_grad=True) mask = x > 0 y = x[mask].sum() y.backward() # x.grad不会是所有位置都有值这里x.grad只在mask为True的位置上有梯度,其他位置为None。如果后续代码对这个梯度做别的运算,很容易出问题。
第二,inplace操作要极度小心。x[0] = 1这种原地修改在requires_grad=True时可能导致梯度计算错误,甚至直接报错a leaf Variable that requires grad is being used in an in-place operation。建议不要对叶子张量做任何原地修改。
第三,view和reshape在梯度回传上通常没有问题,因为它们本质上是同一个张量的不同视图,梯度会正确地映射回去。但如果你在一个view的结果上做了原地修改,再对原张量做计算,梯度方向就会乱套。这种情况最好的防御手段是:不修改、不多想;真要修改,先clone()。
5.3 内存拷贝与性能注意事项
维度调整带来的隐秘性能问题主要体现在内存拷贝上。transpose、permute返回非连续张量,contiguous()会触发拷贝;flatten在非连续张量上也可能拷贝;reshape在非连续时同样拷贝。
写高性能推理代码时,尽量让数据从头到尾保持连续。一个实用建议:数据预处理阶段尽量一次性调整好布局(比如直接转成C, H, W),不要在模型里反复permute再contiguous。PyTorch模型内部可以使用ModuleList加自定义的Sequential层来固定布局,这样对整个pipeline的可读性和性能都有帮助。
另一个常见困惑是clone、detach、copy_的区别。clone会复制张量并保留梯度计算图,detach会断开梯度连接但共享内存,copy_是原地拷贝数据。在维度调整场景中,如果你只要一个不含梯度关系的新张量,最安全的方式是:
new_x = x.detach().clone()如果不需要梯度,也可以用x.detach()直接获得new view,再决定要不要clone。这个组合拳在写数据增强、做可视化、存特征时都非常实用。
5.4 写一个minimal案例来验证维度操作
排查维度问题时,我的习惯是从小规模数据开始验证,不直接在大张量上试。比如写这样一个最小案例:
def debug_dim(fn): a = torch.randn(2, 3, 4) b = torch.randn(2, 3, 4) try: out = fn(a, b) print("OK, shape:", out.shape) except Exception as e: print("Error:", e) debug_dim(lambda a, b: torch.cat([a, b], dim=1)) # 结果(2, 6, 4)这样能快速验证某个操作的维度语义,也方便在Stack Overflow或GitHub Issue里描述问题时贴出来,别人一看就懂。毕竟维度调整这种问题,代码比文字解释快得多。
还有一个小技巧:用x.stride()查看张量的内存步幅。很多维度错误通过stride一眼就能看出来,比如做了transpose之后,stride会变成反向的。看到stride和我们预期不一致,就知道当前张量是不连续的,后面该不该调contiguous心里就有数了。这个函数在常规教程里提得不多,但排查维度问题非常有用。
最后再分享一个心得。我在项目里常用的做法是给每个关键张量写一个形状注释,类似:
# x: (batch, time, feature) -> (batch, feature, time) x = x.transpose(1, 2)虽然看起来多打几个字,但调试时能省很多时间。尤其当你拿到别人写的模型代码,满屏都是permute和reshape,没有注释就只能靠猜。形状就是张量的“身份证”,写清楚它,比什么都强。