ARTICLE DETAIL

资讯详情

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

大模型训练显存优化:从峰值测量到预算决策

大模型训练显存优化:从峰值测量到预算决策 模型又报CUDA out of memory的那一刻通常不是坏事而是预算决策没做够。大模型显存优化这条路走下来最难的不是知道权重占了多大而是训练侧把每个 step 里真实的显存峰值测出来再根据这个峰值做预算决策。推理场景相对简单输入过去、输出回来显存水位基本稳定训练侧完全不是这样一个 step 里前向传播把激活值越堆越高反向传播又需要额外空间优化器更新那一下还要动参数和状态数据每一步都是动态起伏的。如果不做测量只按模型参数量乘个系数拍脑袋后面大概率会在某个凌晨被 OOM显存溢出叫醒。这篇文章我会按自己项目里的实操路径来写先讲训练侧显存的构成再给一套可以照抄的测量流程最后说拿到数字之后怎么做预算决策。适合正在做大模型微调、想换更大 batch、纠结要不要上多卡或者考虑 LoRA、QLoRA 方案的团队。不需要读者有很深的基础但涉及命令和代码的部分建议对着自己的环境试一遍。1. 训练侧测量和推理侧为什么完全是两码事很多刚转过来做大模型训练的朋友会先拿推理的经验套训练模型文件多大显存就预留多大再额外加个 20% 的余量。这个思路在推理时部分成立但放到训练侧几乎必翻车。训练比推理多出来的东西最核心的是三块梯度、优化器状态、前向过程累积下来的激活值。这些量在单个 step 内是高度动态的。前向时网络一层层算下去每一层的中间结果都要暂存供反向传播使用到了反向时又要新开辟空间去存算出来的梯度优化器更新那一下还不算最占内存但需要把所有参数的状态数据都装在内存里。想在某个固定时刻用nvidia-smi截一张图去判断“够不够”基本等于盲人摸象。1.1 一个 step 里的显存水位变化远比你想的剧烈我用一个实际的观察来讲这个问题。PyTorch 训练循环里每个 step 大致会经历前向计算激活值达到近期峰值如果 batch 大、序列长这里可能直接撑满。计算 loss额外暂存一小部分标量梯度和图信息。反向传播根据自动求导机制回传显存继续增加如果开了 activation checkpointing这里会重算部分激活峰值会出现在不同位置。优化器更新对参数做 in-place 更新学习率、momentum 等状态常驻显存。梯度清零或释放显存回落。所以我一直建议团队做训练侧测量时不能只打点看均值要看的是max_memory_allocated这种“峰值”指标。峰值才是预算的依据。你哪怕平均只用了 30GB峰值一旦顶到 80GB 卡上依然会炸。1.2 为什么要单独建立一套测量习惯而不是靠搜别人配置大模型的显存占用和具体框架实现高度相关。同样一个 7B 模型你用 Hugging Face Transformers 跑和用原生 PyTorch 手写 layer 跑显存占用能差出好几 GB。数据并行时梯度和优化器状态要不要分片、通信量怎么算、通信 buffer 开多大都会直接影响结果。更关键的是别人的“经验配置”往往没有告诉你背后的 batch、sequence length、混合精度策略、是否开启 gradient checkpointing。这几个变量一换显存能轻松翻倍。我自己踩过最典型的一次坑照搬社区里某个 13B 模型的微调配置对方显卡是 80GB我这边只有 48GB我把 batch 减半以为能跑结果仍然 OOM。后来一查对方开了 activation checkpointing我没开激活值多出 20 多 GB。测量习惯的价值就在这儿不看别人答案拿着自己的脚本实测才知道数字到底从哪来。2. 先把单步训练显存拆成五笔账再谈测量训练侧显存并不是一笔糊涂账。做测量之前最好先把显存花在哪里拆明白这样后续看到报告数据才能快速判断异常出现在哪。分项内容影响因子模型权重包括基础权重、可能存在的额外参数参数量、精度梯度参数对应的梯度参数量、精度优化器状态Adam 的 exp_avg、exp_avg_sq 等参数量、精度、优化器类型前向激活值每一层中间输出、注意力矩阵、复算辅助信息batch size、序列长度、层数、隐藏层大小临时/通信 buffer分布式通信、算子临时内存、PyTorch 缓存分配器余量并行策略、框架实现、是否使用 checkpointing2.1 模型权重、梯度、优化器状态是“静态地基”这三样在训练过程中基本是固定大小所以是最容易预估的。以 7B 模型为例半精度权重7B 个参数每个 2 字节约 14GB。梯度同样是 7B 个 2 字节浮点约 14GB。如果用 Adam 优化器还要额外保存 exp_avg 和 exp_avg_sq 两个状态通常是单精度 fp32每个 4 字节合计 28GB。9 齐之下就是 56GB还没算激活值。所以一个 7B 模型全参数微调理论上底线就在 60GB 上下这时你再看手头 48GB 的卡就该知道不是硬调 batch 能解决的而是策略方向要改变比如换 LoRA 或者用 ZeRO 做参数切分。优化器状态差异也要注意。SGD 优化器几乎不存额外状态Adam 类优化器则要存两份额外状态。同样是 7B全参 SGD 和全参 Adam静态地基相差 28GB。很多人做决策时只听“模型多大”却忘了问一句“优化器是什么”这里就会埋雷。2.2 前向激活值才是训练侧真正的大头激活值指的是前向传播过程中各层输出的中间结果。它不像权重可以事先估算得那么死而是由 batch size、序列长度、层数、Transformer 隐藏维度共同决定的。序列长度的影响尤其夸张注意力分数矩阵的大小是“序列长度 x 序列长度”量级把 2048 提到 8192激活值不是翻 4 倍而是可能膨胀 16 倍起步。我习惯打个比方权重和优化器状态像是仓库里固定的货架而激活值是你现在摊在工作台上的零件摊多少取决于你同时铺开了多少东西。训练时为了反向传播零件必须留到用完才能收。这也就解释了为什么长序列任务里激活值往往会超过权重本身。2.3 通信 buffer、缓存分配器和框架额外开销经常被漏掉PyTorch 默认使用 CUDA caching allocator它申请过的显存不会立刻归还给显卡驱动而是留在自己的缓存池里供后续快速复用。所以你在nvidia-smi里看到的 reserved memory往往比模型实际 allocate 的 used memory 高出不少。这部分在显存紧张的卡上尤其要重视因为哪怕你的模型只用了 44GB缓存放着不还系统依然可能报 OOM。分布式训练里还有 NCCL 通信所需的 buffer。数据并行在反向传播结束时需要 AllReduce 梯度通信量和模型参数量成正比。启用 ZeRO 分片后还有分片状态和 all-gather 的临时缓冲区。这些项目加起来少则几百 MB多则几十 GB测试时绝不能只算权重一项。所以我在做测量时关注两个指标torch.cuda.memory_allocated()实际分配用量。torch.cuda.memory_reserved()包括缓存在内的总量。如果 reserved - allocated 长期偏大说明要么显存复用不充分要么存在大量短生命周期的大内存分配这时候调整 batch size 或者用torch.cuda.empty_cache()去验证才能看到真实水位。3. 一套能直接照抄的训练侧测量流程测量看起来很 trivial但实际操作中很容易出错。这部分我按自己跑项目的顺序写从上到下就是一套流程建议第一次做时老老实实全走一遍。3.1 第一步不吃任何优化开关先跑一个最小 step我先说为什么要单独跑“最小 step”。正式训练脚本里通常加载了数据、checkpoint、断点续训逻辑还有各种信息打印这些都是干扰项。测量阶段我会单独写一个最小化脚本只做加载模型、构造随机输入、跑一次前向和反向、打印峰值显存。import torch import torch.nn as nn from transformers import AutoModelForCausalLM model AutoModelForCausalLM.from_pretrained( your/model-name, torch_dtypetorch.bfloat16, device_mapcuda, ) batch 1 seq_len 512 input_ids torch.randint(0, 32000, (batch, seq_len), devicecuda) torch.cuda.reset_peak_memory_stats() model.train() output model(input_idsinput_ids, labelsinput_ids) loss output.loss loss.backward() peak_allocated torch.cuda.max_memory_allocated() peak_reserved torch.cuda.max_memory_reserved() print(fpeak allocated: {peak_allocated/1024**3:.2f} GB) print(fpeak reserved : {peak_reserved/1024**3:.2f} GB)这个小脚本的价值在于先把 baseline 打出来。比如启动 CUDA context、加载模型、进入第一个 forward 阶段的底噪。之后再逐步加上你想要的条件才能判断差异来自哪里。注意几个细节一定要调reset_peak_memory_stats()否则算出来的可能是历史累计峰值。一定要设置model.train()模式和 eval 差得很远。随机输入就好不需要真实数据因为这一步只测内存不测收敛。3.2 第二步按关键维度做三到五组温标测量基线跑了之后下一步是调节关键维度找到显存随规模变化的曲线。一般我会做几组固定组合batch size 不变序列长度从 512 涨到 2048再看峰值。序列长度不变batch size 从 1 涨到 4、8再看峰值。打开和关闭 gradient checkpointing 各跑一次记录峰值差距。这里要关注的不是算得准不准而是“尺度感”。你做过一轮之后就能得到一个简单的经验规律例如“序列长度每翻一倍峰值大约涨多少 GB”。以后做预算决策时不用再翻旧文档直接拿这个规律推。每跑一组我都建议把结果记成一张小表配置激活/梯度/静态峰值 allocated峰值 reserved是否通过batch1, seq512, 无 ckpt梯度不降低X GBY GB是batch1, seq2048, 无 ckpt序列引起持续扩大X GBY GB否batch1, seq2048, 开 ckpt通过时间换空间X GBY GB是这比笼统地说“用 checkpointing 能省内存”要有效得多。3.3 第三步在正式训练循环里做持续监控而不是只靠外部命令很多团队会在宿主机上用nvidia-smi -l 1去持续刷显存这在我看只能算是“急诊监控”不能作为“预算数据”。因为训练进程内部的内存释放和申请非常快一秒级别的采样间隔可能刚好错过几个真正的高峰。更好的做法是把监控埋进训练循环里每隔固定步数就记录一次峰值顺便和 step 对账。核心思路是每个 step 开始前重置峰值统计step 结束后读取峰值再记录下来。def compute_peak_memory(): torch.cuda.reset_peak_memory_stats() return None # 这只是一个初始化占位 for step, batch in enumerate(train_dataloader): if step % 20 0: torch.cuda.reset_peak_memory_stats() input_ids batch[input_ids].cuda() output model(input_idsinput_ids, labelsinput_ids) loss output.loss loss.backward() optimizer.step() optimizer.zero_grad() if step % 20 0: peak torch.cuda.max_memory_allocated() / 1024**3 reserved torch.cuda.max_memory_reserved() / 1024**3 print(fstep {step}: peak allocated{peak:.2f}GB, reserved{reserved:.2f}GB)实测下来正式训练里的峰值一般出现在优化器 step 前后几分钟还有可能是通信和激活重算叠加的位置。记录一时的峰值比事后跑nvidia-smi复盘要准得多。3.4 第四步补充分布式和通信上下文测量如果你的目标是多卡训练单卡 measurement 还不够。多卡场景里显存构成多了通信 buffer而且梯度同步会让峰值走向更复杂。实践中最稳的做法是先单卡跑出基础峰值。再以数据并行方式启动两个进程同时观察每张卡的峰值。如果发现单卡和多卡峰值有显著差异重点检查 NCCL buffer 和 ZeRO 策略是否开启了碎片化。另外我强烈建议在分布式场景打印 GPU 显存的“均匀性”。数据并行下如果某张卡明显偏高往往说明数据切分不均或者模型并行没做对齐这个靠肉眼刷nvidia-smi很难盯住但训练日志里每卡峰值一拉出来就能看出来。4. 拿到实测峰值之后预算决策具体怎么做测量只是手段真正的目的是决定“这套训练配置该在哪张卡上跑、 batch 该多大、并行策略该用什么”。这部分我总结成一套可执行的预算决策流程。4.1 先算静态底盘和动态蓄水池拿到峰值数据后我会先把训练侧的显存分成两层静态底盘模型权重 梯度 优化器状态 常驻通信 buffer。动态蓄水池活跃激活值 临时算子缓冲区 缓存分配器余量。预算决策的基本判断就是静态底盘 动态峰值 显卡可用显存同时所有并行实例的累计不超过每卡上限。公式本身没什么高深关键在“动态峰值”的取值。你测出来的峰值就是决策基础宁可高估一点也不要卡着边界。因为我踩过太多次的坑是显存刚好跑过验证集却在第 80 个 step 因为一次特殊长度的 batch 崩掉。4.2 单卡能装直接跑但给预留余量单卡方案我喜欢按总量的 10% 到 15% 作为预留。比如一张 48GB 的卡实际可用也就 45GB 上下因为 CUDA context 和驱动本身要占一部分。如果测量峰值显示 42GB那大概率能跑如果显示 44GB哪怕没报错我也会考虑降低 batch 或开启 checkpointing而不会赌它“刚好不炸”。4.3 单卡装不下先调材料再调并行策略顺序很重要很多人一上来就想开 ZeRO、开多卡其实顺序反了。我建议按下面的链路依次判断降低动态峰值。优先减 batch size、减序列长度或者用 gradient checkpointing 换空间。降低静态地盘。切到 LoRA 或 QLoRA只训练少量参数或者换优化器比如从 Adam 换成 8-bit Adam。开启模型并行或分片。单卡实在装不下的情况下再做张量并行、流水线并行或 ZeRO。为什么这个顺序重要因为并行策略会带来额外的通信开销和显存冗余在还没压榨单卡空间时直接上多卡反而浪费硬件。我见过不止一次一个 13B 模型在 48GB 单卡上调整好激活就能老实跑完微调根本不需要开多卡。几个常用手段的显存收益大致如下手段节省方向代价减 batch size激活值可能影响收敛稳定性需要用梯度累积补回减序列长度注意力激活值可能影响长距离依赖开启 gradient checkpointing激活值训练时间增加约 20%-40%使用 LoRA/QLoRA优化器状态和梯度模型表达能力和收敛范围受限开启 ZeRO 分片参数、梯度、优化器状态通信量增加需要多卡4.4 预算决策在什么时候要重新测量经验里最容易犯的错是改了一个小参数觉得对显存影响不大就不做重测。事实是gradient_checkpointing开关一换峰值位置会变优化器类型一换静态底盘直接新增几十 GB数据并行从 2 卡变成 4 卡通信 buffer 也会翻。任何会影响权重、梯度、优化器状态或激活值任何一个环节的改动都应该重新走一遍测量流程不用全套至少跑一次最小 step 看峰值。5. 微调场景里的预算差异LoRA、QLoRA、全参微调怎么选大模型微调是目前训练侧最常碰到的场景。很多团队手里没有特别大的显存又想跑 7B、13B 甚至更大的模型于是 LoRA 这类方案就成了热门选择。但不同微调方案的显存结构差异非常大预算决策里的“静态底盘”完全不一样。5.1 全参微调和 LoRA 的显存地图几乎不在一个数量级全参微调时静态底盘包含完整权重、完整梯度、完整优化器状态。这是我们在第 2 节算过的 56GB 等级。而 LoRA 训练时原始权重处于冻结状态虽然前向传播还是整模型参与但梯度只算 LoRA 适配器那一小部分优化器状态也就只对应那部分参数。所以同样一个 13B 模型全参微调随便轻松吃到 80GB激活值还没往上叠就快爆了。LoRA 的静态底盘小得多可能剩出 50GB 以上的空间给激活值、长序列、大 batch 使用。代价是 LoRA 不训练大部分底层参数模型最终表达能力可能受限。预算决策不能只看显存还得看任务效果。我的做法通常是用 LoRA 先跑通 baseline快速拿到可用的效果同时把显存预算压低确实需要全参微调时再单独开一轮训练用更多卡分摊成本。5.2 QLoRA 和 4-bit 量化基座激活值反而成为关键瓶颈QLoRA 更进一步把底座模型量化成 4-bit 权重进一步压低静态底盘。这样能跑更大模型但训练时的激活值依然是以半精度或更高精度存储的模型越大激活越可能成为瓶颈。所以 QLoRA 并不是万能解药它解决的是“静态底盘过大”的问题激活和序列长度该占的还是占。这类场景的预算要点看静态模型量化后权重从 fp16 的 2 字节参数收窄到 0.5 字节4-bit 量化后能大大降低。看动态长序列、大 batch 依然依赖激活空间需要测。看混合如果开gradient_checkpointing激活省了但前向重算时间会增加CPU 端观察到的整体耗时可能会涨。我之前做过一个 33B 模型的 QLoRA 实验静态底盘压到很低单卡预留空间充足实际训练却因为 sequence length 长达 8192 而导致激活值高达 30 多 GB。当时差点怀疑 Quantization 出了 bug后来一查就是注意力矩阵占满了显存。这种事多碰几次就能深刻体会到测量是贯穿整个预算决策的。5.3 微调预算决策里的常见误判这里列几个我在项目里反复看到的误判读者朋友可以对号入座只算模型文件大小不看优化器状态。只看到 LoRA 参数少就认为显存一定小忽略了中间层的激活值和 forward 过程。开了gradient_checkpointing后只看总量下降但没留意峰值位置可能变化。在单卡上验证全参微调 OOM立刻选择上多卡而不是先检查是否可以通过激活重计算解决问题。6. 当实测数据和预算不一致我的排查路线测量和预算毕竟不是一次就准。有时候测出来单卡明明够正式训练却会 OOM这时候不要慌按我下面的路线去定位基本能在半小时内找到问题。6.1 先确认是不是 reserved 和 allocated 的差太大遇到 OOM第一反应不是去调小模型而是看日志里 allocated 和 reserved 的关系。如果 allocated 还有余量reserved 已经顶到上限那说明缓存分配器留下了太多不归还的显存。这个场景里尝试在 loss.backward() 后调用torch.cuda.empty_cache()有时候并不能根治反而可能引入碎片更有效的做法是直接调小 batch或者开启PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True来减少预留段碎片。6.2 再查序列长度是否均匀分布式训练里数据长度不一会造成每个 step 的显存峰值差异很大。比如多数 step 用 seq_len1024偶尔来一条长样本到 2048峰值突然冲高。这种波动靠手工查很难发现建议在数据加载时记录样本长度分布同时观察训练日志里有没有“特定 step 附近 OOM”的规律。6.3 回头看动态图或者内存快照PyTorch 提供了比较详细的内存快照能力遇到复杂问题可以调用torch.cuda.memory._record_memory_history()拍快照再用torch.cuda.memory._dump_snapshot()导出分析。这个工具在定位“某个算子临时申请了巨量显存”的场景特别有用。有个真实的例子某个训练改完 attention 实现后显存飙升快照一看是einsum在构建中间矩阵时一次性创建了巨大临时张量。那个临时张量生命周期极短allocated 峰值一冲而过但正好把上限突破了。这种情况常规的“看模型权重尺寸”完全查不出来只有内存快照能直观看到哪一行代码在冲高。6.4 建立“每次改动后测最小步”的规矩回到文章开头那句话训练侧测量的价值不在测量本身而在形成准入门槛。我的团队规定改任何与模型结构、长度、优化器、并行策略相关的配置必须先跑那套最小 step 测量脚本。跑了之后把峰值贴在训练记录里再决定是否继续。这套规矩执行下来OOM 发生率显著降低而且每个人对显存逻辑的理解也同步加深了。7. 最后分享一点个人习惯我现在做任何大模型训练项目第一件事不是调超参而是先花半小时把显存测量脚本跑明白。即使有现成配置文件我也会快速试一张卡上的峰值确保后续各种决策都基于真实数据而不是道听途说。显存优化这件事预算做得越早训练过程就越顺畅每次省下来的那点折腾时间远比自己之后再翻车值钱。
返回列表