news 2026/10/9 4:17:14

PyTorch张量操作精讲:索引分片、合并与维度调整

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch张量操作精讲:索引分片、合并与维度调整

开头先说一下为什么要把这几个操作单独拿出来写一篇。不管是做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) # 展平成一维,等价于flatten

unsqueeze在深度学习里有个极为经典的场景:单条样本转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 matchcat拼接维度不一致打印两个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 Doubledtype不匹配统一.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,没有注释就只能靠猜。形状就是张量的“身份证”,写清楚它,比什么都强。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/9 4:17:12

Pi 1.0 发布:原生MCP与Durable如何重塑终端编程代理

最近我把手头一个项目的终端编程工作流彻底重做了一遍&#xff0c;核心原因是 Pi 1.0 正式版发布了。这个版本给我的感觉不是小修小补&#xff0c;而是把终端编程代理这个品类往前推了一大步——原生 MCP 支持加上 Pi Durable&#xff0c;前者让 AI 代理能直接接入整个外部工具…

作者头像 李华
网站建设 2026/10/9 4:17:10

Java web超市管理系统课设:数据库设计与Tomcat部署实战

简介&#xff1a;这是一套基于Java Web的超市管理系统完整项目资料&#xff0c;面向计算机相关专业的在校学生、课程设计或毕业设计需求者&#xff0c;以及希望以真实项目练手的Java Web初学者。项目已通过导师评审&#xff0c;答辩成绩达95分&#xff0c;代码经测试可正常运行…

作者头像 李华
网站建设 2026/10/9 4:17:01

Spring Boot冷链监控平台:从需求到答辩的完整设计与实战

1. 项目概述&#xff1a;冷链监控平台到底在做什么如果你正在准备毕业设计&#xff0c;或者刚接触Spring Boot想找一个靠谱的练手项目&#xff0c;冷链监控平台这个题目我建议你认真考虑。它不像纯电商系统那样烂大街&#xff0c;但业务逻辑又足够典型——实时数据采集、阈值告…

作者头像 李华
网站建设 2026/10/9 4:16:52

Stata调用大模型:catllm实现文本分类与主题发现实战

做实证研究的人应该都经历过这样的夜晚&#xff1a;三千条开放题回答摆在面前&#xff0c;每一段都要人工编码&#xff0c;而明天就要交初稿。我那时一边盯着一家企业客服投诉数据&#xff0c;一边在Stata里来回翻看&#xff0c;忍不住去搜“Stata 调用大模型”&#xff0c;然后…

作者头像 李华
网站建设 2026/10/9 4:16:41

二叉树详解:递归遍历、搜索二叉树与运行时错误排查

1. 先理解二叉树&#xff1a;别被名字吓住&#xff0c;它只是“每个节点最多俩孩子”的树很多朋友学到数据结构&#xff0c;第一个卡住的坎往往不是链表&#xff0c;而是二叉树。链表好歹还能靠“穿珠子”的直觉理解&#xff0c;二叉树一说“递归”“左右子树”&#xff0c;脑子…

作者头像 李华
网站建设 2026/10/9 4:16:20

基于Workerman与WebSocket的多端在线客服系统架构实践

做在线客服这套系统&#xff0c;我前后折腾了差不多一个季度。一开始公司给的需求很简单&#xff1a;"客户在网页上找我们咨询&#xff0c;客服能实时回复就行。"后来需求慢慢长成了四个端&#xff1a;PC网页、手机H5、微信小程序、App。中间我纠结过要不要直接买第三…

作者头像 李华