ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

PyTorch张量操作实战:从view、reshape到维度变换避坑指南

PyTorch张量操作实战:从view、reshape到维度变换避坑指南 说实话先不管 PyTorch 官网那一堆花哨的教程也不管你从哪个视频里看到“从零入门 PyTorch”的集数真正决定你能不能把模型跑起来的从来都是张量操作。为什么这么说因为 PyTorch 整个基础框架说白了就是搭在张量Tensor上的一套自动求导系统。你喂给模型的数据是张量模型里的权重是张量梯度也是张量。你要是连张量的维度变换都玩不转哪怕把 YOLO 的源码翻烂了也改不出自己想要的结构。这篇文章就是把“张量操作”这层窗户纸捅破适合刚装完 PyTorch、还分不清 view 和 reshape 有什么区别的初学者也适合那些写模型写了不少但老在维度上报错的半新手。看完你能直接照着敲遇到问题也知道去哪里排查。1. 张量到底是个什么东西1.1 它不只是个“多维数组”很多教程喜欢把张量解释成“多维数组”这句话没错但特别容易误导人。NumPy 里也有多维数组你用 numpy 也能做矩阵乘法、也能做广播、也能做 reshape。那 PyTorch 的张量和 NumPy 的 ndarray 到底差在哪儿差在两件事GPU 加速和自动求导。GPU 加速很好理解你在 CPU 上算一个 512x512 的矩阵乘法可能还行但要是你在训练神经网络前向传播里随便一层就是几百万次浮点运算CPU 直接卡死。张量可以放到 GPU 上计算也就是把数据从内存搬到显存里运算速度能快几十倍。自动求导更是 PyTorch 的看家本事你定义一个张量时需要设置requires_gradTrue之后所有对这个张量的操作都会被记录在一个“计算图”里调用backward()时梯度会自动算出来。你如果用 NumPy就得自己手动实现反向传播写两行就想摔键盘。我见过有些同学用torch.tensor([1, 2, 3])创建了一个张量然后整天疑惑为什么别人的代码里有.cuda()或者device这种东西就是因为他没意识到张量除了“值”本身还带着dtype、device、requires_grad这些隐性属性。这三个属性一旦不匹配就会产生各种莫名其妙的问题。所以学张量操作第一件事不是背函数而是建立这个意识你手里拿的不是一个孤零零的数组而是一个携带着硬件位置、数值类型、梯度状态的数据容器。1.2 动手之前先想清楚三件事我在带新人的时候要求他们在创建任何张量之前先回答三个小问题这个数据要放在 CPU 还是 GPU如果放 GPU那 CPU 上的数据要先.to(device)搬过去。数值类型是 float32 还是 int64模型权重一般都要 float32标签一般用 int64索引一般用 int64。这个张量需要求梯度吗只有需要训练的参数才需要requires_gradTrue别一股脑全开浪费内存不说回头有时候梯度反传到你想不到的地方。别小看这三个问题。很多人训练 RNN 或者 Transformer 的时候报错“Expected floating point type for target with class probabilities”或者RuntimeError: expected scalar type Long but found Float八成就是第二个问题没管好。你不需要把每个类型都背下来但至少要知道torch.float32和torch.long大概对应什么场景出错了才知道往哪个方向排查。我记得我第一次用 PyTorch 训练一个简单的全连接网络时数据是从 DataFrame 里读出来的里面既有整数又有小数直接转成torch.from_numpy(df.values)结果报错说数据类型不支持。当时就是没意识到 NumPy 数组默认可能是 int64但网络参数要求 float32。你光看数值是对的但类型对不上框架一样不惯着你。这个坑说起来低级踩过的人却不少。2. 创建张量的几类实用方法2.1 从列表或者 NumPy 转过来最直接的创建方式就是torch.tensor([1, 2, 3])但要注意torch.tensor是会复制数据的。如果是从 NumPy 转过来的我建议用torch.from_numpy()因为它是共享内存的转换开销小。这句话什么意思就是你改 NumPy 数组的值对应的张量也会变反之亦然。好处是快坏处是容易在你不注意的时候发生数据被意外修改的“灵异事件”。import torch import numpy as np arr np.array([1.0, 2.0, 3.0]) t torch.from_numpy(arr) arr[0] 99.0 print(t) # tensor([99., 2., 3.])看到了吧这个共享机制在实际工程里是把双刃剑。如果数据在预处理阶段还要频繁改动最好用torch.tensor(arr)复制一份避免后续模型训练时数据被无意中覆盖。如果只是想快速把 NumPy 数据送进网络而且确定不会再改原数组用from_numpy是更优解。从 Python 列表创建目前我还是推荐torch.tensor()。如果你用torch.Tensor([1, 2, 3])这种构造器它其实等同于torch.FloatTensor行为上会有一些隐性默认值。torch.tensor()会根据数据自动推断 dtype而torch.Tensor()永远创建 float32这个区别经常让人困惑。我的原则是什么时候都别省那几个字符用torch.tensor()让类型显式一点。2.2 按形状创建zeros、ones、randn 怎么选造测试数据的时候经常需要一个“形状正确”但内容无所谓的张量。这时候torch.zeros()、torch.ones()、torch.randn()、torch.empty()轮流上场。torch.zeros(3, 4)3 行 4 列全部填 0。适合生成掩码、初始化偏置。torch.ones(3, 4)全部填 1。适合做全连接层的 bias 初始化。torch.randn(3, 4)从标准正态分布里采样适合测试模型前向传播。torch.empty(3, 4)分配内存但里面的值是垃圾桶里捡来的没初始化。只有你确定马上会被覆盖时才用。实际写模型时你还要区分torch.rand_like(t)和torch.randn_like(t)。前者生成 0 到 1 之间的均匀分布随机数后者生成标准正态随机数它们都会保持和t相同的形状。这两个函数在做 Dropout 或噪声注入的测试时特别方便不用手动传 shape。下面是一段肉眼可见的实例import torch a torch.zeros(2, 3) b torch.ones(2, 3) c torch.randn(2, 3) d torch.full((2, 3), 7.5) # 全部填 7.5 print(a device:, a.device) print(d shape:, d.shape)torch.full((2, 3), 7.5)可能很多人不知道它用来生成“全填充为指定值”的张量比torch.ones() * 7.5直观得多而且不会引入额外的乘法运算节点。2.3 dtype、device、requires_grad 的坑这三兄弟是真的能联手坑人。我举个最典型的例子x torch.randint(0, 10, (3,)) y x.float() # 转 float32 z torch.randn_like(y) z.requires_grad_(True) # 原地设置 requires_grad第一行是 int64第二行转成 float32第三行用randn_like保证了形状和 dtype 都和 y 一致第四行原地开启梯度记录。这个流程在代码里看起来平平无奇但如果你没有第二行直接对 int 型张量计算梯度PyTorch 会直接告诉你“RuntimeError: Only Tensors of floating point and complex dtype can require gradients”。这也是我在实际问答里见到概率极高的一类报错。还有一个隐性坑torch.arange(0, 10)生成的默认 dtype 在某些版本里是 int64但torch.range(0, 10)是 float32。很多人用range以为它和 Python 的range一样结果类型对不上。现在 PyTorch 官方其实都建议用torch.arange替代torch.range因为torch.range的行为太反直觉它会包含结尾值而且 dtype 默认不是整数。记住一个小原则索引类张量用 int64计算类张量用 float32别混着用。3. 形状操作“变形记”view、reshape、permute、transpose3.1 view 和 reshape 的区别到底在哪这是我在各技术论坛上被人问得最多的问题之一。一句话解释view只对“内存连续”的张量有效它直接复用底层内存不复制数据reshape更聪明如果张量内存连续它就和view一样不会复制如果不连续它会先拷贝一份让内存连续再改变视图。那“内存连续”是什么意思你可以把张量的内存想象成一排连续的小格子。普通创建一个 3x4 张量数据在内存里就是按行依次排下来的连续。但是你一旦执行了t.t()转置逻辑上得到的是一个 4x3 张量可底层内存顺序还是原来的 12 个格子没法按新形状直接“顺序解读”这时候内存就不连续了。view看到这种情况直接报错reshape则会先contiguous()拷贝出连续的内存再返回新视图。x torch.arange(12).reshape(3, 4) y x.t() # transpose try: z y.view(4, 3) except RuntimeError as e: print(view 报错:, e) z y.reshape(4, 3) print(z)所以我个人建议新手优先用reshape因为它更智能不容易在形状变换时栽跟头。想真正提升性能或者对内存布局很敏感的人再去深挖view和contiguous()。遇到“viewsize is not compatible”这类报错别慌改成reshape或者先调用.contiguous()再view一般就通了。3.2 permute 和 transpose 千万别混淆x.transpose(dim0, dim1)只能交换两个维度例如二维矩阵转置就是x.transpose(0, 1)。x.permute(dims)能一次性把所有维度重新排列比如x.permute(2, 0, 1)表示把原来第三维挪到最前面。这两个操作都是视图操作不复制数据所以改数据会互相影响这点要记住。尤其在图像处理里PyTorch 的默认张量布局是[C, H, W]也就是通道在前。你用 OpenCV 读出来的图片是[H, W, C]也就是通道在最后。想塞进 PyTorch 模型就需要把通道维度换到最前面。一开始我见过有人写img.permute(2, 0, 1)后来也有人用img.transpose(0, 2)再transpose(1, 2)结果自己也绕晕了。我建议只记permute(2, 0, 1)这一个写法图像从[H, W, C]换成[C, H, W]它是最明确的。permute 还有一个非常容易踩的坑变换之后张量通常不连续。你如果接着调用view十有八九又报“view size is not compatible”。正确的姿势是先permute再.contiguous()再view。img torch.randn(224, 224, 3) # 模拟 HWC img_chw img.permute(2, 0, 1) # 变成 CHW print(img_chw.shape) # torch.Size([3, 224, 224])3.3 unsqueeze 和 squeeze 的实战场景这两个函数名字长得像兄弟但功能正好相反。unsqueeze是在指定位置插入一个新的维度长度为 1squeeze是把长度为 1 的维度去掉。很多人刚开始根本不知道为什么要“多加一个长度为 1 的维度”。我给你说一个最常见的场景。假设你有一条数据形状是[4]表示一个样本的 4 个特征。但 PyTorch 的线性层nn.Linear要求输入至少是二维的第一维是 batch size第二维是特征数也就是[1, 4]。所以你得用x.unsqueeze(0)把它变成[1, 4]。如果最后你想把这一个样本的特征向量恢复成[4]就用x.squeeze(0)。还有卷积层输入要求[B, C, H, W]。如果你做单张图片推理图片读进来是[C, H, W]也要先unsqueeze(0)变成[1, C, H, W]。这就叫“手动补一个 batch 维度”。这里要注意.squeeze()默认是去掉所有长度为 1 的维度但你可以给它传参数只去除指定位置。比如x.squeeze(dim2)如果你原以为第二维长度是 1但实际不是它不会报错而是什么都不做。这个“静默不报错”的行为有时候反而会掩盖问题。我自己写代码时能显式传dim就尽量显式传避免无意中把多个长度 1 的维度全删了。3.4 用形状推导法避免“维度灾难”张量操作多了以后我发现自己陷入一个怪圈不停地加 reshape、permute、unsqueeze最后把维度搞得连自己都不认识了。后来我养成了一个习惯每一步关键变换都手动写下当前的shape再写下目标shape用笔推一遍。举例来说我有一个 attention 矩阵原始形状是[batch, heads, seq_len, seq_len]我想去掉中间那个 heads 维度把它合并到 batch 里去。我会写# 当前: [B, H, T, T] attn attn.permute(0, 2, 3, 1) # 变成 [B, T, T, H] attn attn.reshape(B, T, T * H) # 合并掉 H如果我不写注释一周后回来看这段代码一定得重新推半天。所以这里强烈建议关键形状变换前后都写注释。调试时也可以临时加print(x.shape)。这个方法笨但特别有效尤其是刚开始学 Transformer 的人Shape 推导熟练了各种框架代码读起来都会顺畅很多。4. 张量的索引、切片与组合拼接4.1 像 NumPy 一样索引但小心维度套路PyTorch 的索引语法和 NumPy 高度一致x[0]取第一个x[:, 1]取所有行的第二列x[1:, :2]切片x[[0, 2]]按列表取行。这些我都默认你会真正想提醒的是下面两点。第一布尔掩码索引返回的是一维张量。比如x[x 0]你得到长度等于“满足条件的元素个数”的向量原来的形状全丢了。这在做损失计算时有用但你要是想保留矩阵结构需要另想办法。第二索引结果一般情况下会共享数据也就是返回的是视图。什么意思你修改y x[0]中的元素x也会变。如果你不想影响原张量就要用y x[0].clone()。第三花式索引fancy indexing有时候会复制数据比如x[[0, 1, 2]]这种索引操作和切片不一样它不保持内存连续性。这个对普通算法影响不大但如果你做自定义算子就会碰上性能问题。x torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) mask x 4 print(x[mask]) # 一维tensor([5, 6, 7, 8, 9]) print(x[x[:, 1] 2]) # 按某一列过滤行第2列2的所有行4.2 torch.cat 还是 torch.stack这是个问题torch.cat沿现有的某个维度拼接要求除拼接维度外其余维度完全一致。torch.stack是增加一个新的维度把所有张量“堆”起来。你可以这么理解cat是把几块积木沿长度方向接起来stack是把它们叠成几层。举例a torch.tensor([[1, 2], [3, 4]]) b torch.tensor([[5, 6], [7, 8]]) c torch.cat([a, b], dim0) # 形状 [4, 2] d torch.cat([a, b], dim1) # 形状 [2, 4] s torch.stack([a, b], dim0) # 形状 [2, 2, 2]从形状结果就能看出区别。cat不增加总维度数stack总是增加一个维度。数据加载的时候如果你想把多个特征矩阵拼在一起用cat如果你想把一批单独的样本“摞”成批次张量用stack。实际写 dataloader 的时候我最常用的组合是stack配合unsqueeze。比如 batch 里的每条样本是[seq_len, hidden]我想得到[batch, seq_len, hidden]直接torch.stack(batch_list, dim0)就行代码短而且不容易出错。4.3 split 和 chunk 不是一回事torch.split可以按“每块大小”拆也可以按“每块数量”拆接口比较灵活。torch.chunk则是固定拆成 n 块如果张量长度不能被 n 整除最后一块会小一点。这个“整除”问题很容易埋雷。比如你有一个长度为 10 的张量想用torch.chunk(x, 3, dim0)拆成 3 块PyTorch 给你返回 3 个张量长度分别是 4、4、2不是均匀的。如果你用torch.split(x, 3, dim0)则返回长度为 3、3、3、1 的多个张量。初学者常常忽略这个差别导致循环里处理每一块时以为大小一样结果最后一块尺寸不一样就崩了。x torch.arange(10) for piece in torch.split(x, 3): print(piece.shape) # torch.Size([3]) # torch.Size([3]) # torch.Size([3]) # torch.Size([1]) for piece in torch.chunk(x, 3): print(piece.shape) # torch.Size([4]) # torch.Size([4]) # torch.Size([2])如果你希望风格稳定就用split因为它可以指定固定尺寸。如果你只是想把一个 batch 平均分到多张卡上但还不能保证一定能整除那么chunk的“最后一块比较小”行为也许更符合预期但一定记得处理好边界条件。5. 常见形状错误与排查经验5.1 我最常碰到的三类张量报错这些报错几乎每个人都会遇到我把它们整理成一个速查表希望能帮你省点时间。报错信息常见原因解决方向view size is not compatible操作时张量内存不连续或新形状元素总量与原形状不一致用reshape代替或先.contiguous()再viewExpected scalar type Long but found Float/ 反之dtype 不匹配如标签用 float 而损失函数期望 long用.long()或.float()转换mat1 and mat2 shapes cannot be multiplied矩阵乘法中 inner dimensions 不一致打印两个张量的shape检查维度对应关系Expected all tensors to be on the same device一部分张量在 CPU一部分在 GPU统一调用.to(device)Sizes of tensors must match except in dimensiontorch.cat或广播时除指定维度外其他维度不一致打印参与拼接的每个张量的shape对齐形状这些报错本身不要怕怕的是你不看报错信息就瞎改。我的建议是遇到报错第一件事把print(x.shape, y.shape)加在报错那一行之前看清楚两个张量的形状再动手。大多数情况下错误自己就能定位了。5.2 用几行代码定位维度问题我在调试程序时会习惯性地写一个极小的“形状检查脚本”把关键的张量挨个打印出来。你用不着用什么高级调试器print 大法最直接。def describe(name, tensor): print(f{name}: shape{tuple(tensor.shape)}, dtype{tensor.dtype}, device{tensor.device}) a torch.randn(2, 3) b torch.randn(3, 4) describe(a, a) describe(b, b) try: c torch.matmul(a, b) describe(c, c) except RuntimeError as e: print(matmul error:, e)这样一眼就能看出a是(2, 3)b是(3, 4)乘法的最后维度是 3 对 3没问题。如果写成a是(2, 3)b是(2, 4)你马上就能发现内维度 3 和 2 不匹配。5.3 错误排查的三个独门心得第一不要盲目用.cpu()和.numpy()。很多初学者为了打印张量先.cpu().numpy()结果遇到“can‘t convert cuda:0 device type tensor to numpy”这种报错。其实你直接print(tensor)就行张量支持直接输出不非得转 NumPy。第二当你怀疑模型训练出了问题先检查损失函数的输入。分类问题最常见的就是预测值没经过softmax之前是全连接层输出形状是[batch, num_classes]标签形状是[batch]。这两个形状对应清楚损失函数的坑基本就没了。第三requires_grad也是排查时的重要指标。如果某个张量本来需要梯度但你在中间做了一次detach()或者从 NumPy 转过来时没设置requires_gradTrue反向传播到这里梯度就断了。症状往往是“模型不更新”或“某些层梯度为 None”。排查方法很简单在loss.backward()之前打印一下模型参数的.grad看看是不是 None。for name, param in model.named_parameters(): if param.grad is None: print(f{name} 没有梯度)写在最后我的一些小经验张量操作这种东西看十篇教程不如自己敲一遍。刚开始练的时候我也会把view、reshape、permute混着用现在反而更克制了能不用高级花活就不用代码越直白越好。我给自己的规矩是所有跨维度的变换一律写注释标注变换前后的形状所有拼接操作先打印参与拼接张量的形状所有涉及requires_grad的地方尽量显式声明而不是靠隐式继承。最后分享一个我一直在用的小技巧遇到任何张量形状问题先别急着搜网上的“报错解决方案”花一分钟在小纸上把(batch, seq_len, hidden)这类中间形状推一遍。很多时候不是 PyTorch 函数有问题而是你自己把维度关系搞拧了。把形状推清楚再看代码其实就顺了。这也是我见过的大多数 PyTorch 高手不约而同的习惯。
返回列表