ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

机器学习张量详解:从形状报错到显存爆炸的排查思路

机器学习张量详解:从形状报错到显存爆炸的排查思路 张量这个词学机器学习的同学几乎天天见但真正理解它并不是背下“张量就是多维数组”这一句话就够的。我接触过不少刚开始跑深度学习教程的人代码里到处都是tensor.shape、torch.randn可一遇到报错还是分不清是数据形状不对、维度顺序反了还是数据类型不一致。这篇就围绕“机器学习里的张量到底是什么、代码里怎么判断它的状态、报错后按什么顺序排查、显存为什么会突然爆掉”这几个实际问题展开。适合刚入门机器学习、正在跟 PyTorch / TensorFlow 教程走、或者期末复习想把概念串起来的人。看完之后你至少能自己定位大部分张量相关的报错而不是对着屏幕猜。1. 先搞清楚张量在机器学习里到底是什么1.1 从标量、向量、矩阵到张量一条递进理解路径很多人第一次接触张量是在线性代数或者深度学习课程里。那时老师会说“张量就是多维数组”听起来和 numpy 里的 ndarray 没什么区别。其实先这样理解没错但还不够。从简单到复杂递进标量一个单独的数字比如3.14它是 0 阶张量。向量一维有序数字集合比如[1, 2, 3]它是 1 阶张量。矩阵二维表格比如[[1,2],[3,4]]它是 2 阶张量。张量三阶、四阶乃至更高维的数组统称。比如一批彩色图片[32, 256, 256, 3]就是一个 4 阶张量。所以“阶数”可以简单理解成“有几个维度”。tensor.ndim在 PyTorch 里返回的就是阶数tensor.shape返回的是每个维度的大小。我在给别人讲这个概念时最常用的一个类比是“俄罗斯套盒”。标量是最里面那个数字往外每套一层就多一个维度。矩阵是两层的盒子三阶张量是三层的盒子。机器学习里的数据经常要套很多层因为要同时装下样本数量、通道数量、图片高度、图片宽度这些信息。1.2 为什么框架不直接用 ndarray非要叫张量这是很多初学者最容易困惑的点。既然 numpy 也有多维数组PyTorch / TensorFlow 为什么还要自己搞一套 Tensor 类型核心区别不是“维度更多”而是张量带上了机器学习需要的“计算能力”。同样是存储一批数字ndarray 主要负责数值计算但 Tensor 在它的基础上还支持自动求导能记住数据是怎么一步步计算出来的反向传播时可以直接算梯度设备管理可以搬到 GPU 显存里运行也可以放回 CPU计算图构建训练过程中的每个操作都会进入一张动态或静态计算图作为模型参数、模型输入、模型输出的统一数据格式。所以你在代码里看到的model(x)x是 Tensor模型权重也是 Tensor损失值还是 Tensor。整个训练链路完全被张量贯穿。这里需要澄清一个重要误解数学里定义的“张量”有更严格的含义涉及多重线性映射、坐标变换规则。机器学习工程里说的“张量”大多数时候就是一个带梯度、带设备信息的多维数组。初学者不用在数学定义上先死磕先按“能跑计算的多维数组”理解后面再回头看数学定义会轻松很多。2. 判断一个张量是否合理先看四个属性2.1 shape形状决定一切张量最值得关注的就是形状。神经网络每一层本质上都是在做“张量到张量”的变换层与层能不能衔接完全看形状对不对。举个例子。一个 batch 里有 32 张彩色图片每张图片高 256、宽 256有三个颜色通道。那么输入张量的形状通常写成[32, 3, 256, 256]通道在前NCHW或者[32, 256, 256, 3]通道在后NHWC。不同框架有不同默认顺序PyTorch 常用 NCHWTensorFlow/Keras 早期格式里 NHWC 也很常见。文本数据也是一样。一个 batch 有 32 个句子每个句子长度是 40那么 token 索引张量形状是[32, 40]。如果再接一个词嵌入层把每个 token 映射成 128 维向量就会变成[32, 40, 128]。为什么形状这么重要因为矩阵乘法、卷积、注意力计算都依赖特定的维度布局线性层要求输入最后一维等于in_features卷积层对输入输出通道数量有明确要求多头注意力里 Q、K、V 的特征维度必须一致。所以代码里最常出现的一行调试语句就是print(tensor.shape)你不知道数据现在长什么样就不知道模型为什么会报错。2.2 dtype类型不一致是隐藏炸弹除了形状dtype 也直接影响训练是否正常。常见类型有float32深度学习默认主力精度和内存占用比较平衡float64更精确但更占内存速度通常更慢int64标签和索引常用uint8原始图像像素常见范围是 0 到 255bfloat16/float16混合精度训练常用。图片数据读出来经常是uint8但神经网络里的运算通常要求float32。如果不做转换模型可能会隐式帮你转也可能直接报类型错误。我的建议是在数据预处理阶段就显式转好不要依赖框架自动处理。标签类型也要注意。交叉熵损失函数在很多框架里默认要求标签是整数张量比如int64。如果你手滑把标签转成float32轻则报错重则训练过程完全不符合预期。2.3 device数据必须和模型在同一位置训练时会经常遇到类似报错Expected all tensors to be on the same device意思是模型在 GPU 上但输入数据还在 CPU 上。解决办法是input_tensor input_tensor.to(device) model model.to(device)学习阶段完全可以用 CPU 跑通小任务不要一上来就追求 GPU。先确保逻辑正确再把数据搬到显存里加速。很多新手一上来就配 CUDA结果环境没配好反而被设备问题干扰连基本流程都跑不顺。2.4 梯度和计算图训练和推理看到的张量不一样训练模式下参与计算的张量会记录操作历史形成一个计算图。requires_gradTrue的张量在有梯度需求的操作链上会保留中间结果反向传播时才能求出梯度。模型参数本身就是张量它默认需要梯度。输入数据是否需要梯度取决于你的设计。如果是训练通常要求输入数据不需要梯度但权重需要如果是纯推理可以用torch.no_grad()包住计算不做梯度追踪省内存也省时间。有一个细节我建议新手现在就养成习惯看损失值时用loss.item()把它转成 Python 数字而不是直接print(loss)。直接打印会同时输出计算图信息很啰嗦也容易误导人。import torch x torch.randn(2, 3) y x.sum() print(x.shape) # torch.Size([2, 3]) print(x.dtype) # torch.float32 print(x.device) # cpu print(y.shape) # torch.Size([])0 阶张量也就是标量这段代码几乎是入门阶段最常用的“望闻问切”方式。3. 把张量流动过程完整走一遍3.1 数据管线输出的就是张量很多人对张量的理解停留在“数据读入之后才变成张量”其实数据管线从读取到预处理的每一步都在决定最终张量形状。以图像为例。你用 PIL 或 OpenCV 读出来的可能是一个 HWC 结构的数组像素值还是 0 到 255 的整数。要做的是转成浮点张量、调整维度顺序、归一化、加上 batch 维度最后才能送进模型。中间少一步后面都容易出问题。以文本为例。原始字符串不能直接参与矩阵运算要先通过分词器转成 token 索引序列每个 token 对应词表里的一个整数。之后再查嵌入表得到一个[序列长度, 嵌入维度]的矩阵。再加上 batch 维度变成[batch, 序列长度, 嵌入维度]。表格类数据相对简单通常是[样本数, 特征数]比如 1000 个样本、20 个特征就是[1000, 20]。这里最容易忽略的是“坐标顺序”。同样是图片有的框架用[B, C, H, W]有的用[B, H, W, C]。顺序错了模型不一定报错但结果会非常离谱。所以数据进入模型之前先打印一行print(batch.shape)确认。3.2 模型中间层到处是张量变换模型的前向过程就是一连串张量变换线性层输入(B, in_features)权重(in_features, out_features)输出(B, out_features)卷积层输入(B, C_in, H, W)输出(B, C_out, H_out, W_out)注意力层输入(B, T, C)通过不同投影矩阵得到 Q、K、V最后输出(B, T, C)。所以排查模型问题时一个很实用的做法是把模型里每一层的输入输出形状都打印出来对照网络定义检查。不要凭感觉猜。我一般会先把 batch size 设为 2甚至 1用一个随机张量输入模型看能不能跑通前向。如果前向都过不了问题基本出在维度衔接上。model MyModel() dummy torch.randn(2, 3, 224, 224) out model(dummy) print(out.shape)这个“随机张量试跑”的方法能帮你把“数据问题”和“模型问题”分开。随机张量能跑通说明模型结构没大问题再回头查真实数据格式。3.3 loss 到 backward为什么损失通常是一个标量张量训练时损失函数返回的通常是一个 0 阶张量也就是标量。比如对一批样本的交叉熵取均值变成一个数。为什么要标量因为反向传播需要从最终损失开始把梯度一层层传回去。如果 loss 是一个向量梯度就会变成雅可比矩阵计算复杂也很容易出问题。所以代码里经常看到loss loss.mean() loss.backward()做过几次实验后你会发现用loss.item()记录训练日志非常方便它拿到的是纯 Python float不参与计算图不会干扰梯度。4. 内存占用与“对齐”改个参数就爆显存是怎么回事4.1 先学会估算张量占用显存不够是机器学习最经典的报错尤其在你把 batch size 调大、分辨率调高之后。很多人第一反应是把 batch 调小但如果你能提前估算张量占用排查速度会快很多。占用的基本公式张量内存 元素个数 × 每个元素字节数float32每个元素占 4 字节float64占 8 字节int64占 8 字节uint8占 1 字节。举个例子。一个(224, 224, 3)的 RGB 图片如果是float32元素数是 150528内存大约等于 150528 × 4 602112 字节约等于 588KB。一个 batch 32 张输入本身大约 19MB。听起来不大但深度学习真正吃显存的地方往往是中间层的特征图以及反向传播时需要保留的中间结果。如果一个模型有几十层卷积每层都会输出多个 feature map中间张量会比输入数据大很多倍。所以你会看到输入只是 224×224 的图片训练时显存却能轻松超过好几个 GB。估算的意义不是要求你手工算准每个字节而是建立“元素个数 × 类型字节数”这个意识。当你把输入从 224×224 改成 512×512 时某些层的内存占用不是线性增长而是平方级增长因为宽高同时变了。4.2 “内存与张量对齐”到底是啥热词里出现了“内存与张量对齐”这个词看起来有点底层但它其实会直接影响代码运行效率甚至导致报错。先分清两个层级的“对齐”。底层的内存对齐是指分配内存时让起始地址按 2、4、8、16 等字节对齐CPU 或 GPU 读取效率更高。这是框架和操作系统层面处理的普通用户一般不用管。但用户层经常遇到的是张量的“连续性”也就是 contiguous 属性。当你对一个张量做转置、维度交换、切片后它在底层可能不再是一块连续内存。打印is_contiguous()会得到False。这时如果调用view()这类要求内存连续的 API就可能报错。实际经验是如果你只是对张量做transpose/permute之后直接送入需要连续内存的操作可能会碰到报错。解决办法是加一个.contiguous()把数据复制成一块连续内存。x torch.randn(2, 3, 4) y x.permute(0, 2, 1) # 内存可能不连续 z y.contiguous() # 复制出一块连续内存但不要什么操作都加.contiguous()。它是复制操作会带来额外内存和耗时。只有报错明确提到 non-contiguous或者你需要调用view/ 某些要求连续内存的 API 时再加。还有一个常见操作是reshape。某些框架里reshape会自动处理连续性问题而view不会。所以新手阶段如果只是要改变形状优先用reshape更省心如果需要严格控制底层存储方式再区分view和reshape。4.3 显存或内存不够时的常规调整顺序如果你训练时遇到 OOM也就是 out of memory不要盲目堆机器配置。先按这个顺序排查缩小 batch size这是最直接、最有效的办法降低输入分辨率或序列长度但要注意评估对精度的影响检查是不是在保留整个计算图而不是每轮只取数值推理阶段确认开启了torch.no_grad()尝试混合精度训练比如把参数和激活值部分使用 float16但需要框架支持清理不用的中间变量必要时显式删除大张量。低配机器能跑通 demo不代表它能跑完整训练。我一般会用最小配置先验证代码逻辑然后再逐渐加大 batch 和分辨率找到当前机器能承受的上限。如果你只是学习默认配置通常够用。如果要长期跑实验就得提前规划数据批大小、日志保存和输出目录毕竟张量并不只是“理论概念”它真实地占着你的内存和显存。5. 张量相关报错的排查链路5.1 典型报错到底长什么样见过足够多报错后你会发现很多错误信息是“同一类问题换了个说法”。常见的包括size mismatch两个张量进行运算时形状对不上Expected input to have X dimensions, but got ...维度数不对很常见于把 3 维张量传给要求 4 维输入的层mat1 and mat2 shapes cannot be multiplied矩阵乘法时维度不匹配Expected all tensors to be on the same device设备不一致Expected tensor to be double but got floatdtype 不一致non-contiguous内存不连续。第一次看到这些英文报错时别慌绝大多数都能在报错倒数几行里看到具体形状。5.2 按顺序排查先复现再打印再改参数我自己的排查顺序一直很固定因为乱试参数只会让情况更乱先跑一个最小样例比如 batch size 1或者直接用随机张量打印数据集的 batch 形状确认输入格式打印每个关键层前后的 shape、dtype、device对照模型定义检查输入维度是否满足每一层的要求确认 loss 是标量再看反向传播是否报错如果报错提到连续性检查是否在transpose/permute之后用了view最后才考虑框架版本、依赖版本、模型权重加载这类外部因素。为什么先复现因为最小样例能快速判断是“模型结构问题”还是“真实数据问题”。如果最小样例能跑真实数据却报错那大概率是数据维度、类型或路径的问题。5.3 一个典型误判看着像“功能不支持”其实是维度顺序反了这类问题我遇到太多次了。你以为模型不支持某种格式的输入其实只是维度顺序和框架约定不一致。比如模型期望输入是(B, T, C)而你喂给它的数据是(T, B, C)它可能会报出维度不匹配的错。又比如图像模型期望(B, C, H, W)你的数据是(B, H, W, C)形状上可能碰巧能过一部分层但结果全乱。修正方式是用permute调整维度顺序而不是用reshape。这两者区别很大reshape是按当前逻辑顺序重新排布元素它不关心维度的“含义”permute是交换维度的位置元素之间的关系会保持对应。举个例子一个(B, H, W, C)的张量要变成(B, C, H, W)应该用permute(0, 3, 1, 2)而不是reshape(B, C, H, W)。用 reshape 很容易把像素数据彻底打乱造成“训练能跑但损失不下降”这种更隐蔽的问题。6. 学习建议从认识张量到独立跑通一个项目6.1 入门阶段别急着追大模型现在各种大模型、多模态模型满天飞很多新手一上来就想跑大模型。但我建议先把最基础的小任务跑通比如手写数字分类、房价预测、简单的文本分类。原因很简单小任务里你能非常直观地看到输入张量怎么流经每一层梯度从 loss 一路回传权重张量怎么更新。大模型动辄几十亿参数一旦报错根本不知道是数据处理问题、模型结构问题还是显存和并行策略问题。我第一次跑通一个简单的全连接网络时做的事很简单在每个关键步骤后打印张量形状从输入一直看到输出和 loss。这个过程远比背概念更有效。6.2 起步阶段建议掌握的 6 个张量操作如果你刚入门不需要背几十个 API。先把下面这些操作用好操作作用典型使用场景reshape改变张量形状扁平化特征图、拼接后调整尺寸permute/transpose交换维度顺序通道维度和序列维度调整unsqueeze/squeeze增加或删除长度为 1 的维度给单样本加 batch 维度cat/stack拼接张量合并多个特征、构造 batchexpand/repeat复制扩展维度广播计算、构造掩码to(device)切换设备CPU 与 GPU 数据迁移这些操作会覆盖大多数入门场景。等真正需要处理生成式模型、多模态模型时再学更复杂的索引和高级 API 也不迟。6.3 期末复习和面试准备时怎么看“张量”如果是为了准备考试张量通常会从三个角度考概念定义什么是 0 阶、1 阶、2 阶、N 阶张量形状判断给你一段代码或一个数据写出shape数学性质涉及到线性变换时张量的坐标变换规则。如果是面试重点通常会放在实践上。比如让你解释为什么loss必须是标量、为什么模型输入要归一化、为什么view和reshape不一样。抓着这几个问题准备要比死记定义有用得多。6.4 关于学习资源的一句话建议经常有人问机器学习入门书买谁的。周志华《机器学习》、李航《统计学习方法》、吴恩达课程、李宏毅课程这些都是大家常提的资源。我的看法是如果目标是考试和理论优先看教材把符号和定义搞懂如果目标是动手做项目先跟着代码教程走遇到问题再回查教材千万不要从头到尾硬啃一本教材却不写一行代码。理论和代码必须交叉学。张量这个概念尤其如此你只在纸上画再多次也不如真正在代码里打印一次tensor.shape来得深刻。最后留几个自查点踩过几次坑之后我发现很多张量相关问题不是工具能力不够而是前置环境和输入材料没有处理干净。每次遇到莫名其妙的现象我都会回头确认这几件事数据形状是否和模型预期一致数据 dtype 是否是模型需要的类型数据和模型是否在同一个设备上是否在transpose/permute之后错误使用了view显存或者内存占用是否已经超过当前机器上限训练时是否误开了梯度追踪导致计算图越攒越大。把这几个点变成你的下意识反应张量相关报错的解决速度会明显变快。建议你现在就打开一个深度学习项目在训练脚本里提前加上 shape、dtype、device 的打印。跑一遍你会对张量有完全不一样的感觉。
返回列表