news 2026/10/9 3:35:38

PyTorch张量操作进阶:索引分片、合并与维度调整完全指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch张量操作进阶:索引分片、合并与维度调整完全指南

刚开始用PyTorch跑模型的时候,我大部分时间都在跟报错搏斗,其中出现频率最高的不是loss爆炸,而是各种和shape有关的提示。后来才慢慢发现,与其说是模型结构复杂,不如说是在张量的索引分片、合并和维度调整这些基本功上不够熟练。这批操作几乎每天都要用,而且一旦理解透,很多网络结构代码看起来就像拼积木一样简单。这篇学习笔记围绕这三个方向做一次系统梳理,把关键原理、常用写法和容易踩的坑一起记录下来,适合刚接触PyTorch的初学者,也适合想查漏补缺的日常使用者。

先说结论:张量操作不难,难的是脑子里要有“形状动态变化”的画面。下面的内容不按API文档顺序罗列,而是从最底层的形状理解开始,逐步过渡到索引分片、合并和维度调整,最后用一个小案例把几个操作串起来。

1. 张量的形状:理解维度的第一步

1.1 从shape看张量的嵌套结构

很多新手看到torch.randn(2, 3, 4)这种写法就发蒙,不知道生成了一个什么东西。其实可以把张量的shape想象成快递箱的嵌套结构:最外层有几个箱子,每个箱子里有几行,每行里有几列。比如torch.randn(2, 3, 4),就先数外层2个“大箱”,每个大箱里有3个“中盒”,每个中盒里有4个“小格子”。这样生成的张量,shape就是[2, 3, 4],维度ndim是3,元素总个数numel是2*3*4=24。

import torch x = torch.randn(2, 3, 4) print(x.shape) # torch.Size([2, 3, 4]) print(x.ndim) # 3 print(x.numel()) # 24

理解这个嵌套结构有一个非常实际的好处:当你看到一个报错提示Expected 4D input时,你就得先自查当前张量到底是几维。我曾经把一个[64, 1, 28, 28]的张量当成[64, 28, 28]用,结果卷积层报维度错误,定位了很久才发现是自己默认少写了一维。所以在任何操作之前先打印shape,是成本最低的排错手段。

1.2 维度的顺序为什么不能搞反

张量的维度顺序不只是数字上的不同,它对应了数据的物理含义。以图像为例,PyTorch里默认布局是[N, C, H, W],即batch数量、通道数、高度、宽度。如果哪次不小心把某个操作写成了[N, H, W, C],后面再接卷积或者归一化层时,数据含义就全错了,但它不会立刻报错,而是会在模型里“带病运行”,结果就是训练半天精度上不去。

再比如自然语言处理里的常见形状[batch, seq_len, hidden],第一个维度是样本数,第二个是序列长度,第三个是每个时间步的特征维度。调整维度时经常会把seq_len和hidden弄混,最后全连接层接收的输入不符合预期。我的习惯是在代码里给关键张量写清楚注释,例如“当前形状[B, T, D],B是batch,T是序列长度,D是hidden size”,这样后续做transpose或permute时不容易搞反。

1.3 先用一行代码掌握张量基本属性

在实际写代码前,我建议先建立一个“三件套”意识:看到一个新张量,立刻打印shape、dtype和device。这三个属性决定了一个张量能不能参与后续运算,也决定了你是否需要转换。

y = torch.zeros(3, 5) print(y.shape) # torch.Size([3, 5]) print(y.dtype) # torch.float32 print(y.device) # cpu

dtype不一致会导致自动类型转换失效,比如一个float32和一个float64张量相加偶尔会报错。device不一致则是GPU环境下最常见的报错来源。虽然这篇笔记主要讲索引分片和维度调整,但保持打印张量属性这个习惯,能让你在排查所有张量问题时都更快定位原因。

2. 索引与分片:精准取数和切片

2.1 基本索引:维度之间用逗号分隔

PyTorch张量的索引逻辑和Python列表很接近,区别在于多维张量需要用逗号把不同维度的索引分隔开。一维张量就像普通列表,支持正数索引和负数索引,负数表示从末尾开始数。

a = torch.arange(10) print(a[0]) # tensor(0) print(a[-1]) # tensor(9) print(a[5:8]) # tensor([5, 6, 7])

二维矩阵用a[i, j]取某一行某一列。注意这里不能写成a[i][j],虽然它也能工作,但那是连续两次索引,可读性差且更容易出错。用逗号分隔才是PyTorch的风格。

m = torch.arange(12).reshape(3, 4) print(m[1, 2]) # tensor(6),第二行第三列 print(m[0]) # tensor([0, 1, 2, 3]),第一行 print(m[:, 1]) # tensor([1, 5, 9]),所有行的第二列

把索引落到“嵌套结构”上看就很直观:第一个下标选外层箱子的编号,第二个下标选中盒的编号,第三个下标选小格子的编号。想取某个维度里的全部内容就用冒号:代替。

2.2 切片操作:左闭右开和步长

切片是索引的扩展,格式沿用了Python的start:stop:step,而且区间是左闭右开,也就是说1:5包含下标1、2、3、4,不包含5。这个规则几乎所有Python使用者都知道,但在多维张量里依然经常出问题,尤其是step为负数时,含义会反过来。

t = torch.arange(20).reshape(4, 5) print(t[1:3, :]) # 取第1行到第2行,所有列 print(t[:, ::2]) # 所有行,列方向每隔2个取一个 print(t[::-1]) # 行方向倒序

实际写数据预处理时,我最常用的是在时间维度上切片,比如取某个序列的前半段,或者从第10帧取到第20帧。这里有一个容易忽视的细节:切片返回的是视图,而不是拷贝。视图意味着原始数据存储被复用,你修改切片结果,原张量也会跟着变。比如sub = t[0]; sub.fill_(0)会真的把原张量的第一行全部改成0。如果不想影响原数据,记得用.clone()一下。

2.3 布尔索引:按条件筛选数据

除了按下标取数,PyTorch还支持用布尔张量做掩码筛选。这是处理分类标签、过滤异常值时非常好用的功能。你先用一个比较表达式得到形状相同的布尔张量,然后把它作为索引放进方括号里,得到的是所有满足条件的一维元素。

scores = torch.tensor([0.8, 0.3, 0.6, 0.9, 0.2]) mask = scores > 0.5 print(mask) # tensor([ True, False, True, True, False]) print(scores[mask]) # tensor([0.8000, 0.6000, 0.9000])

不只是简单比较,布尔索引可以用&、|、~组合多个条件,注意一定要加括号,否则会报语法错误。

scores = torch.tensor([0.8, 0.3, 0.6, 0.9, 0.2]) filtered = scores[(scores > 0.5) & (scores < 0.9)] print(filtered) # tensor([0.8000, 0.6000])

布尔索引在样本筛选上非常直观:一张特征图[N, C, H, W],配合一个标签向量labels,可以用feature[labels == 1]把所有类别为1的样本一次性取出来。注意这里的labels长度必须等于第一个维度的长度,否则无法对齐。

2.4 索引结果到底是视图还是新张量

这是初学PyTorch最容易掉进去的坑。笼统地说,普通连续切片和单个下标索引返回的是原张量的视图,内存共享;而使用布尔掩码、整数列表或张量做高级索引时,通常会返回新拷贝,不再共享内存。

origin = torch.arange(12).reshape(3, 4) sub_view = origin[1:3] sub_copy = origin[[1, 2]] sub_view.fill_(999) print(origin[1]) # 会看到999,因为origin被改了 sub_copy.fill_(0) print(origin[1]) # 没有变化

判断是否共享内存,可以用torch.is_shared(),或者直接做一次修改实验。对我来说,保险策略是:只要下一步操作可能修改取出来的数据,先问一句“原始数据还能不能再动”,如果不能确定,就统一.clone()一次,代价是额外的内存占用,但能避免很多隐蔽的逻辑错误。

3. 合并:cat与stack的使用辨析

3.1 cat():在已有维度上拼接

torch.cat把多个张量沿着某个既有维度拼起来,相当于把几摞数据直接粘在一起。最常用的场景是把多个batch的数据合并成一个更大的batch,或者把不同来源的特征图堆叠到通道维上。

a = torch.randn(3, 4) b = torch.randn(5, 4) c = torch.cat([a, b], dim=0) print(c.shape) # torch.Size([8, 4]) d = torch.randn(3, 2) e = torch.cat([a, d], dim=1) print(e.shape) # torch.Size([3, 6])

使用cat有一个硬性要求:除了拼接的那个维度,其他维度的形状必须完全一致。比如a是[3, 4],想和[5, 4]在dim=0拼接没问题,因为第二维都是4;但如果另一个是[5, 3],就必须先把自己的形状调整成[5, 4]再来。这个规则看着简单,遇到多维度张量时还是容易眼花,我在代码里一般会先把涉及拼接的张量形状打出来,逐维对齐一遍。

3.2 stack():创建一个新维度来堆叠

与cat不同,torch.stack不要求在已有维度上对齐,而是把所有张量当成独立对象,在指定位置新增一个维度,然后把这些对象装进这个新维度里。最典型的场景是:你有一批形状相同的向量,想把它变成一个矩阵,每个向量单独占一行。

v1 = torch.tensor([1, 2, 3]) v2 = torch.tensor([4, 5, 6]) s0 = torch.stack([v1, v2], dim=0) s1 = torch.stack([v1, v2], dim=1) print(s0.shape) # torch.Size([2, 3]) print(s1.shape) # torch.Size([3, 2])

从结果能直观看到dim=0是把两个向量按行堆叠,变成两行三列;dim=1是沿列方向堆叠,变成三行两列。注意stack要求所有张量形状一模一样,不像cat只要求非拼接维度一致。如果你需要给一组标量增加一个“样本”位置,stack比cat更合适。

3.3 合并时最容易踩的坑

第一个常见坑是张量类型不一致。比如一个float32张量,另一个是int64张量,cat大概率报错,因为PyTorch不会自动帮你做类型转换。解决办法是先统一为同一个dtype,a.float()或b.long()。

第二个坑是设备不一致。一个在GPU上,一个在CPU上,合并时也会报错。尤其是长时间写代码后,两个张量来源不同,设备可能悄悄变了。处理方式是先.to(device)统一。

第三个坑是合并后丢失了维度语义。例如两个[batch, seq_len]的张量用cat(dim=1)合并,得到[batch, seq_len1 + seq_len2],看起来好像是在拼接时间步。但如果两个张量代表不同序列,你应该用stack得到[batch, 2, seq_len]作为新的批次维度。语义混淆在后续处理中很难查,我建议在想清楚“这是在延长已有维度,还是引入新的维度”之后再选择cat还是stack。

3.4 神经网络里合并张量的典型场景

实际模型代码里,合并操作大多出现在特征融合和数据加载阶段。比如多模态模型中,图像特征和文本特征分别经过各自网络后,需要把两路特征拼到一起,一般用cat在通道维或特征维上完成。

img_feat = torch.randn(16, 512) text_feat = torch.randn(16, 256) combined = torch.cat([img_feat, text_feat], dim=1) print(combined.shape) # torch.Size([16, 768])

再比如训练循环里,你需要把一个epoch的所有batch特征收集起来做可视化或测试,可以用list先存起来,最后torch.cat拼成一个大张量。这种情况下要注意,列表里的张量必须都来自同一个设备,且除了拼接维度外,其他维度形状一致。还有一个小技巧:如果列表里的张量形状不一致,可以用stack前先做padding或裁剪。

4. 维度调整:把张量掰成需要的样子

4.1 view()与reshape():连续内存的区别

view是PyTorch里最直接的调整形状方法,但它有一个前提:张量在内存里必须是连续的。所谓连续,可以理解成数据在底层存储中按“从左到右依次填充”的顺序排列。大多数新建的张量是连续的,但一旦做了转置、permute等操作,内存就不再连续,此时直接view会报错。

x = torch.randn(2, 3, 4) y = x.permute(0, 2, 1) # 转置后不连续 # print(y.view(2, 12)) # RuntimeError: view size is not compatible print(y.reshape(2, 12)) # 正常

reshape更宽容:能复用视图就复用,不能复用时就拷贝数据生成新张量。所以日常代码里,如果只是想变形状,优先用reshape更省心;如果明确知道张量是连续的,或者希望结果与原始数据共享内存,用view更高效。不过我个人的习惯是,在数据不需要共享内存的场景下直接写reshape,少想一层。

4.2 permute()与transpose():换轴的正确方式

transpose用于交换两个维度,比如把[batch, seq_len, hidden]变成[batch, hidden, seq_len],用transpose(1, 2)。permute则更灵活,可以一次把所有维度重排,参数是维度的新顺序。

x = torch.randn(2, 3, 4) p1 = x.permute(2, 0, 1) print(p1.shape) # torch.Size([4, 2, 3]) t1 = x.transpose(1, 2) print(t1.shape) # torch.Size([2, 4, 3])

permute传入的是目标形状中各位置对应原维度的索引。例如x原本是[2, 3, 4],想变成[4, 2, 3],就意味着新张量的第0维来自原第2维,第1维来自原第0维,第2维来自原第1维,因此要写permute(2, 0, 1)。这个规则一开始容易绕,我的记忆方法是:把permute的参数看成一个映射表,位置上的数字表示“新维度从原来的第几个维度来”。

转置操作返回的通常是视图,而且会让内存不连续。如果转置之后还要做一些小批量操作,最好先.contiguous()再view或reshape。这里要注意contiguous()不是改变数据内容,而是重新申请内存并按照当前逻辑顺序复制一份。

4.3 squeeze()与unsqueeze():增删长度为1的维度

squeeze删除长度为1的维度,unsqueeze增加一个长度为1的维度。这两个操作特别适合在做广播计算和模型输入对齐时使用。

x = torch.randn(3, 1, 4) print(x.squeeze().shape) # torch.Size([3, 4]),所有长度为1的维都被去掉 print(x.squeeze(1).shape) # torch.Size([3, 4]),指定去掉第1维 print(x.unsqueeze(0).shape) # torch.Size([1, 3, 1, 4]) y = torch.randn(5) print(y.unsqueeze(-1).shape) # torch.Size([5, 1])

为什么需要这种操作?最常见的原因是让张量参与广播计算。比如一个形状为[batch, 1]的偏置向量想加到[batch, feature]的特征上,如果某些维度对不上,就需要先用unsqueeze或squeeze把形状调整为[batch, 1]或者[1, feature],让广播规则能正确生效。另一个场景是模型输入规定必须是4维,但你的数据只有3维,这时在某个维度上用unsqueeze补齐即可。

4.4 expand()与repeat():广播式扩展与物理复制

expand和repeat都可以让张量在某个维度上变大,但两者有本质区别。expand不分配新内存,它只是把原张量在逻辑上重复显示,相当于“广播”;repeat会真正复制数据,生成新的内存。这个区别对内存占用和后续修改行为影响很大。

base = torch.tensor([[1], [2], [3]]) # [3, 1] e = base.expand(3, 4) # [3, 4],逻辑上重复 r = base.repeat(1, 4) # [3, 4],物理复制 print(e) print(r)

expand时,想让某一维变大,这一维原来必须是1,比如[3, 1]可以expand成[3, 4],但不能把[3, 2]硬扩成[3, 4],那属于猜测数据,PyTorch会直接报错。repeat没有这个限制,参数表示每个维度要重复几次,比如repeat(1, 4)表示第一维重复1次,第二维重复4次。

实际使用中,expand适合处理偏置、均值等需要广播到大批量上的张量,节省内存;repeat适合你需要真实复制数据并可能独立修改的场景。我曾在自注意力模块里用expand生成掩码,因为操作人数不多,内存没压力,但如果数据量很大,expand的优势就很明显。

5. 综合案例:分组筛选、合并、变形一次打通

5.1 场景设定:从模型输出里筛选两类样本

假设现在有一批图像特征和一个分类器输出的概率分布。特征是形状[N, 4, 8, 8]的四维张量,对应的标签是形状[N]的一维张量,每个标签是0或1。我们需要做三件事:把所有标签为0的样本和标签为1的样本从特征中分别抽出来,然后把两组特征合并成一个新的batch,再将这个batch调整成适合送入全连接层的二维形状。这几乎是图像分类任务调试里最常见的片段。

为了减小示例的规模,设N=6,通道数、高度、宽度都取小值。

feat = torch.randn(6, 4, 8, 8) labels = torch.tensor([0, 1, 1, 0, 1, 0]) mask0 = labels == 0 mask1 = labels == 1 cls0 = feat[mask0] cls1 = feat[mask1] print(cls0.shape) # torch.Size([3, 4, 8, 8]) print(cls1.shape) # torch.Size([3, 4, 8, 8])

这里布尔索引直接把满足条件的样本取了出来,第一个维度从6变成了各自的数量。如果两组数量不对称,也是正常的,只要最后能拼到一起就行。

5.2 合并和变形:让维度重新对齐

接下来用cat把两组特征拼起来,因为是按batch纬度拼,类别的顺序会丢失,所以如果想区分,后续可能要再准备一个拼接后的标签向量。这里只展示特征操作。

combined = torch.cat([cls0, cls1], dim=0) print(combined.shape) # torch.Size([6, 4, 8, 8])

随后要把每个样本的所有通道和空间位置展平,变成[6, 4*8*8]的形状。这里可以使用reshape,因为它不要求张量连续。

flattened = combined.reshape(combined.size(0), -1) print(flattened.shape) # torch.Size([6, 256])

如果后面的模型要求输入形状是[6, 1, 256],还需要用unsqueeze(1)增加一个维度。

final = flattened.unsqueeze(1) print(final.shape) # torch.Size([6, 1, 256])

注意,reshape(combined.size(0), -1)里的-1表示让PyTorch自动推导这个维度的大小。combined总元素数是6*4*8*8=1536,第一维是6,那么第二维自动就是1536//6=256。用-1看起来方便,但前提是你必须清楚总元素个数能整除。

5.3 常见错误与排查技巧速查

踩坑多了之后,我把最常见的几个错误总结成一个速查表,每次遇到报错就对号入座:

错误现象大概率原因处理方式
cat时shape不匹配非拼接维度没有对齐打印两个张量shape逐维检查
stack时维度不一致某张量多了一个长度为1的维统一用squeeze或reshape处理
view报RuntimeError张量内存不连续改用reshape或先.contiguous()
布尔索引后维度丢了mask是一维布尔张量确认筛选方向是否符合预期,必要时先unsqueeze
GPU张量和CPU张量合并报错设备不一致统一用.to(device)
维度调整后结果看起来是乱的permute参数顺序理解错了在草稿纸上标出“新维来自旧维”再写代码

另外一个高效排查技巧是“分段打印形状”。遇到复杂张量操作时,不要想着一步到位,而是每操作一步就打印一次shape,这样能把出错位置精确到某一小段代码。比如在cat之前分别打印两个待拼接张量的形状,一眼就能看出问题。

5.4 一个小习惯:给张量操作加注释

最后一个建议可能与具体API无关,但对长期维护代码特别重要。张量操作链条一长,半年后再打开自己的代码,光看表达式根本想不起每一步在做什么。我一般会在一系列操作的中间变量上直接注释目标形状,甚至在临时变量名里体现形状信息。

# 当前形状: [N, C, H, W] -> [N, L],其中L = C*H*W flattened = combined.reshape(combined.size(0), -1) # 当前形状: [N, L] -> [N, 1, L],适配序列模型的输入 final = flattened.unsqueeze(1)

这不算什么高深技巧,但配合前面说的打印形状习惯,能让我在切换项目或换环境之后,快速找回当时写代码的思路。回头研究PyTorch源码的时候还会发现,索引分片、合并和维度调整这些操作背后其实是同一套张量存储和内存布局的原理,理解得越深,踩坑越少。

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

R语言生态数据分析实战:从群落矩阵到多样性排序绘图全流程

做生态学野外调查的人都知道&#xff0c;从样方里数完物种、称完生物量、记录完环境因子的那一刻&#xff0c;真正头疼的工作才刚刚开始&#xff1a;一摞物种丰度表&#xff0c;怎么变成能写进论文的统计结果和发表级图件。R语言在这条链路上几乎是绕不开的选择&#xff0c;veg…

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

Hadoop MapReduce实战:从CLASSPATH配置到jar运行完整的避坑指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/9 3:32:29

基于大数据的二手房价预测系统:从数据管道到Qt展示的实战解析

做这个基于大数据的二手房价预测系统之前&#xff0c;有个朋友问我&#xff0c;为什么同样的地段、差不多的面积&#xff0c;挂牌价能差出一倍。他说中介靠的是感觉&#xff0c;我说感觉背后应该有数据。于是我用三个月时间把这个系统完整搭了起来——从公开挂牌数据的采集清洗…

作者头像 李华
网站建设 2026/10/9 3:32:05

Linux父进程等待机制详解:僵尸进程回收与wait/waitpid实战

如果你管理的 Linux 服务器上突然冒出一堆状态为 Z 的进程&#xff0c;多半不是系统中毒&#xff0c;而是某个父进程没有做好它应尽的义务。Linux 里的"父进程等待"&#xff0c;从来都不是一句空洞的 sleep&#xff0c;而是一整套围绕进程终结、状态回收、资源释放的…

作者头像 李华
网站建设 2026/10/9 3:31:36

AI赋能一人公司:AI工作流与传统技艺的双引擎实践

深夜十一点&#xff0c;我刚刚把AI生成的户型改造说明改完&#xff0c;客户在微信里发来三个大拇指。电脑左边是我反复调校过的AI工作流面板&#xff0c;右边是一把胡桃木托盘&#xff0c;还停在那天手工打磨到一半的状态。一个用AI跑流程、一个用手艺找意义——这在两年前听起…

作者头像 李华