ARTICLE DETAIL

资讯详情

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

梯度累积:大模型训练显存不够时的关键技巧与实战指南

梯度累积:大模型训练显存不够时的关键技巧与实战指南 跑过大模型训练的人大概率碰到过这样的场面一开训练就OOM被迫把batch_size从32砍到8loss曲线抖得像心电图。这种时候老手通常都会说试试梯度累积gradient accumulation。不少刚接触这个概念的读者会把它理解成“把batch变大”的平替但实际用起来门道很多——什么时候累积、累积多少步、学习率要不要跟着改、多卡怎么配合、AMP下怎么做才不出错这些问题不搞清楚照抄代码换来的可能只是更慢的收敛和一堆莫名其妙的loss spike。这篇文章打算把梯度累积从原理到实战完整拆一遍。内容包括显存瓶颈到底卡在哪、梯度累积的数学等价关系、与学习率和优化器的搭配规则、分布式训练和混合精度的正确姿势以及我这些年踩过的具体坑。适合正在做LLM微调、LoRA、扩散模型训练或者想彻底搞懂训练框架里gradient_accumulation_steps这个参数到底在干什么的读者。1. 为什么需要梯度累积显存瓶颈与batch size的死结1.1 显存到底被谁吃掉了很多人有个惯性认知显存不够是因为模型参数太大。模型参数确实占显存但在大模型训练里真正逼你“压batch”的往往是另一块开销——前向过程的激活值activation。一次常规训练显存大致分为四块占用项说明典型大小模型参数fp16/bf16下约等于参数量×2字节7B模型约14GB梯度与模型参数同形通常同精度7B模型约14GB优化器状态AdamW要存fp32的momentum和variance且每参数两份7B模型约56GB起步激活值每层前向输出与batch_size、序列长度、隐藏层维度强相关随batch线性增长前向激活不是“存一次就完事”的。反向传播需要从最后一层反推梯度每一层的输入都要保留下来供链式法则使用所以一个transformer的激活显存大致正比于 \(layers \times batch \times seq_len \times hidden_size\)。这就是为什么把batch_size从16调到32显存会肉眼可见地往上跳而参数大小反而纹丝不动。我早期训练一个参数量并不算大的BERT-large时batch_size从8改成16直接爆显存看日志才知道是activation那部分在作祟。当时卡上32GB前向激活占了接近40%。1.2 小batch带来的训练问题既然显存装不下大batch那退一步用小batch不就行了也不是完全不行问题是代价明显。小batch下每个batch的梯度估计方差很大一句话描述就是“你每一步走的方向都不太稳”。典型表现是loss曲线方差大、收敛慢甚至优化器在局部震荡出不来。另一个常见副作用是BatchNormBN的均值和方差是“当前batch”内的统计量小batch下统计量噪声极大训练和推理之间的分布差距也会被放大。所以要同时满足两个条件一是显存只够跑小batch二是希望梯度估计尽量接近大batch。梯度累积正是为这个矛盾设计的。2. 梯度累积到底在累积什么数学原理与最小实现2.1 一句话说清核心逻辑梯度累积的操作听起来很简单先不更新参数让模型基于当前参数跑K个micro batch每步都正常做前向和反向把梯度累加到一起然后凑满K步后统一更新一次优化器。这里有一个关键点在这K个micro batch期间模型参数是不变的。如果参数变了后面的梯度就是基于新参数算出来的累积就没有意义了。所以梯度累积的循环里optimizer.step()必须放在“第K步之后”而不能每个micro batch都调用。2.2 数学推导它与一个整体大batch等价吗我们看看SGD的更新式。设一个训练batch里共有B个样本损失函数为 \(L\)则标准mini-batch SGD计算的是平均梯度[ g \frac{1}{B}\sum_{i1}^{B} \nabla \ell_i ]然后更新参数[ \theta_{t1} \theta_t - \eta g ]梯度累积的本质是把这B个样本拆成K份每份m个样本先算每份的梯度 \(g_k \frac{1}{m}\sum_{i1}^{m} \nabla \ell_{k,i}\)再把K个梯度做平均[ G \frac{1}{K}\sum_{k1}^{K} g_k \frac{1}{K \cdot m}\sum_{i1}^{K \cdot m} \nabla \ell_i \frac{1}{B}\sum_{i1}^{B} \nabla \ell_i ]只要每个micro batch的数据互不重叠、且总样本数等于原batch size累积出来的梯度在数学上就等价于直接用整个大batch算出的梯度。这给了我们一个基础结论梯度累积可以视为“用时间换显存”的大batch训练。不过在实操中这种等价性是有条件的。常见的不等价来源有三个BatchNorm的均值方差按micro batch计算每个micro batch前向时dropout的随机mask不同数据增强会改变样本本身。后面第5章再展开说。2.3 最小可用实现PyTorch示例最朴素的PyTorch梯度累积代码相当短accum_steps 4 optimizer.zero_grad() for step, (inputs, labels) in enumerate(dataloader): loss model(inputs, labels) / accum_steps loss.backward() if (step 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()为什么要除以accum_steps因为如果不除K个micro batch的梯度加起来的幅度等于平均梯度的K倍更新步长相当于被放大K倍很容易造成loss直接飞掉。除以K后累积梯度的量级和单次大batch的梯度保持一致。还有一个很多人忽略的细节模型内部如果在每个batch的loss里加了L2正则项除以accum_steps会让L2项也被平均这其实是正确行为。但如果用的是optimizer里的weight decay参数比如AdamW的weight_decay它是在optimizer.step()时才计算的不是在backward时进去的所以累积K步后weight decay的实际频率是原来的1/K。这个行为和大batch训练是一致的不用额外处理但要清楚它和L2正则的语义并不完全一样。2.4 等效有效batch size的计算多卡环境下有效batch基本遵循这个关系effective_batch_size local_batch_size × accumulation_steps × world_size其中world_size是并行卡数。单卡场景直接去掉world_size即可。这组参数是联动的调整任何一个都要重新审视学习率和warmup策略。3. 梯度累积下的学习率缩放那些容易翻车的参数3.1 batch size变了学习率为什么要跟着动很多人踩过这个坑把有效batch从32扩到128学习率原封不动跑出来的效果不仅没提升反而更差。原因不在梯度累积本身而是大batch有一个天然特征——梯度方差更小。想象你在山脊上往下走小batch的梯度方向像喝了酒东倒西歪但综合来看还能下山大batch的梯度方向大多集中在真正下坡的方向但平稳不代表总能走对步子大了容易冲过头。所以batch变大后学习率通常需要相应调大才能维持和大batch匹配的有效步长。3.2 线性缩放规则的适用范围业界最广为人知的规则是Goyal等人提出的线性缩放minibatch size扩大K倍学习率同步扩大K倍。实践论文《Accurate, Large Minibatch SGD》里ImageNet训练把batch从256扩到8192时lr从0.1线性扩到3.2配合warmup在6个epoch内达到不亚于小batch的精度。但注意这是针对SGD的规则不是普适的。线性缩放背后假设是小batch的梯度方差主导了更新噪声而大batch的梯度接近真实梯度。但Adam这类自适应优化器本身有逐元素的学习率归一化机制对整体学习率的敏感度远低于SGD直接按K倍放大反而容易让训练早期出现不稳定的尖峰。3.3 不同优化器的调法参考我自己的经验是分场景处理下面是一个可以照抄的参考表优化器累积K倍后的学习率调整说明SGD/Momentum SGD线性按K倍调整前提是配足够长的warmupAdam/AdamW不调或最多乘以sqrt(K)自适应机制已抵消部分方差变化常用做法是保持lr不变LoRA微调保持lr不变有效batch翻倍对PEFT影响较小动lr反而容易破坏原本调好的基座带BN的CNN SGD按K倍线性调但必须同步加大batch观察BN统计量只靠累积不一定能得到大batch的BN效果举一个具体例子我微调LLaMA类模型时单卡batch2accum_steps8有效batch16学习率从1e-4直接调到1.5e-4训练稳定但如果用SGD微调CNN分类器batch从32变成128学习率我会直接从0.01提到0.04等启动后再看验证集波动。还有一个配套习惯有效batch增大后warmup步数也应拉长。线性缩放理论里热身的目的是让逐渐变大的学习率不把刚初始化的参数推离正常区域大batch 大lr时warmup显得更重要。常用经验值是warmup占总训练步数的3%到10%当K比较大时往10%靠。3.4 累积步数过大时的一套安全做法如果你的accum_steps在8以上或者有效batch翻了好几倍我的建议是一步一步来先固定学习率跑50步看loss趋势如果loss不降或震荡明显再把lr乘以1.2、1.5这样小幅上调不要一步到位。实践中累积步数越大的训练对学习率的微小差异越敏感宁可保守一点。4. 多卡训练里的梯度累积no_sync与allreduce的博弈4.1 单卡到多卡问题立刻变复杂单卡做梯度累积前面那段代码就够了。多卡数据并行时分两种情况你用谈笑风生的方式跑虽然结果通常没大错但这K步里的K次通信带来的开销是实打实的。数据并行多卡下每张卡拿到的是同一个模型的不同数据分片默认DDP在每个backward结束后会做一次梯度allreduce把各卡的梯度求和并广播这样才能保证所有卡上的模型版本一致。如果你在累积循环里老老实实跑了K个micro batchDDP就会同步K次梯度。对于大模型来说allreduce的通信量是模型参数量级别K越大浪费在同步上的时间越多。这就和并发场景下用梯度累积省通信的初衷背道而驰了。4.2 正确姿势只在最后一次backward触发同步PyTorch DDP提供no_sync()上下文管理器可以在前K-1个micro batch里关闭自动梯度同步只在最后一个micro batch结束后触发allreduce。实现起来很自然from contextlib import nullcontext accum_steps 8 model DDP(model) optimizer.zero_grad() for step, batch in enumerate(dataloader): # 前K-1步不做跨卡同步最后一步同步 ctx model.no_sync() if (step 1) % accum_steps ! 0 else nullcontext() with ctx: loss model(batch) / accum_steps loss.backward() if (step 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()为什么这样可以因为在累积阶段各卡的梯度只存在于自己的显存里加上no_sync之后DDP不会在backward时触发allreduce。等到第K步结束时所有卡的梯度都攒够了再做一次allreduce就能得到“全局平均梯度”然后用这个梯度同时更新所有卡上的参数。此时可以确认一下全局有效batch的公式有效batch 单卡micro batch × accumulation_steps × 卡数所以当你设local_batch1、accum8、8张卡时一次参数更新实际吃掉的就是64个样本的梯度信息。这也是多卡训练里非常常见的配置。4.3 多卡边界的三个隐藏坑第一个坑是整除。如果你的数据总量不能被卡数乘accum_steps整除最后几个不足K步的micro batch会搞乱循环某几张卡凑不足K次backward导致DDP在步数边界上hang住。稳妥做法是drop_last把尾部长尾直接扔了或者在末尾手动判断。第二个坑是学习率视角。多卡下累积步数和卡数都放大有效batch学习率要不要调取决于你原本“有效batch”设计成多少而不是只看累积步数。同样是accum88卡和单卡的学习率策略很可能不同。第三个坑是通信频率与ZeRO的配合。DeepSpeed的ZeRO本身会把优化器状态、梯度等切片到不同卡上allreduce和参数更新的逻辑会和PyTorch DDP有所不同。用DeepSpeed时我建议直接用框架提供的gradient_accumulation_steps参数而不是自己改DDP循环否则很容易在ZeRO的梯度划分逻辑上出问题。5. 实战避坑指南AMP、梯度裁剪、BN和日志5.1 混合精度AMP里GradScaler的更新时间混合精度训练时PyTorch用GradScaler把loss放大避免fp16下高频梯度过早下溢到0。梯度累积时最常见的bug是在每个micro batch的backward之后都调用scaler.update()。scaler.update()的作用是根据最近一次梯度过小的比例调整loss scale。如果你在累积中途频繁updatescale值会被不完整的累积梯度带偏轻则训练不稳重则连续出现NaN。正确写法是让scale只在optimizer.step()同频时更新scaler torch.cuda.amp.GradScaler() for step, batch in enumerate(dataloader): with torch.cuda.amp.autocast(): loss model(batch) / accum_steps scaler.scale(loss).backward() if (step 1) % accum_steps 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) scaler.step(optimizer) scaler.update() optimizer.zero_grad()注意上面的unscale_操作它的作用是先把scaler里的scale因子去掉得到真实的梯度再做梯度裁剪。如果忘了unscale就clipclip的对象还是“被放大过”的梯度最终步长会偏差。这句是很多人查半天都没发现的隐藏点。5.2 梯度裁剪到底应该在哪个时机做梯度裁剪的目的是防止梯度过大导致参数瞬间飞出去。采用梯度累积后正确时机只有一个在所有micro batch算完、获得完整累积梯度之后、丢给optimizer之前。如果每个micro batch都做一次clip等于对每一份子梯度分别做截断那累积后的真实梯度可能早已被一次次截断破坏。这在长序列模型里极其致命因为你的有效梯度本身就是跨多个micro batch累计出来的中途截断会丢失大量的方向信息。5.3 数据顺序K个micro batch要像一个batch梯度累积的前提是K个micro batch的数据集合起来等价于一次性“喂进去”的样本。所以你的dataloader循环必须是连续地取K个batch而不是每个micro batch都随机打乱数据。用PyTorch默认DataLoader时shuffleTrue的逻辑是每个epoch开始前打乱一遍全局顺序循环内取到的batch本身是连续的所以梯度累积代码里并不会自动重复抽样。但如果你自己写数据流或者用那种“每个step重新随机抽”的数据接口就得注意累积的K个micro batch必须来自同一批采样策略且互不重叠才能保持等价性。5.4 BatchNorm的统计量不会因为累积而变大这是梯度累积最容易被误解的一点。很多人以为把batch8累积8步就等价于batch64训练。数学上梯度确实接近但BatchNorm的mean和var还是按每个micro batch单独算的——也就是BN仍然面对8个样本的小batch统计噪声问题依然存在。如果你的模型依赖BN且训练集方差大梯度累积并不能解决BN的小batch问题。此时要么尽量调大local batch要么换SyncBN跨卡同步统计量要么在推理时重新估算running stats。不过在纯transformer架构、尤其是LLM和LoRA微调场景里基本用的是LayerNorm而不是BatchNorm所以这个坑主要影响CNN和部分目标检测模型。5.5 日志、scheduler与loss记录累积训练里日志记录也会误导人。你要记录的是每个micro batch的平均loss而不是累积后除以K的“整体loss”。用loss.item()前先把除法拿掉或者记录未除K的原始loss这样日志曲线才不会出现周期性的折痕。scheduler这边通常每个optimizer.step()之后调一次scheduler.step()而不是每个micro step都调。如果你在累积循环内反复推进学习率等价于学习率比预期快了K倍虽然loss不会马上爆但收敛曲线会变得很奇怪。6. 哪些场景真正受益该用与不该用的判断6.1 大模型预训练与LoRA微调的标准配置现在跑LLM和LoRA微调梯度累积几乎是标准配置。LLM的seq_len通常很长激活显存随序列长度和batch显著膨胀一张卡能塞下的micro batch可能只有1或2。想凑够一个像样的batch只能用accum_steps把训练样本累积起来。LoRA场景还有个额外的好处LoRA本身只训练少量低秩矩阵优化器状态占用不大但基础模型的前向激活依然很占显存。所以LoRA和梯度累积天然搭配常见的配置是micro batch1、accum16到64让显存塞满训练、时间换有效batch。6.2 不该硬用梯度累积的情况如果显存充裕直接大batch训练永远是更稳的选择。梯度累积虽然梯度等价但参数更新频率降低、BN统计量偏差、dropout随机性增加这些都会让最终效果和真正的大batch有细微差异。此外在线学习和流式场景不适合累积每条数据都要即时更新模型延迟步进反而有害。还有一些对BN特别敏感的任务累积带来的收益有限最优先方案是同步BN而不是梯度累积。6.3 我的一点个人心得这些年调模型我对梯度累积的经验可以总结成三句话先算清楚有效batch size有多少再决定要不要动学习率多卡必用no_sync单卡也要记得把gradient clipping放到最后一个micro batch之后AMP和累积代码尽量复用已经验证过的模板不要随手改。前阵子调一个7B模型的LoRA单卡micro batch1accum16有效batch16学习率用2e-4没动warmup拉到总步数的5%整体曲线平滑得让我意外。反倒是另一组实验把有效batch翻到128学习率同步调大了50%结果几个step后loss直接飘了。重新把lr降回原始值才恢复稳定。所以我现在的默认操作是累积步数越大学习率调整越克制宁可慢慢来也不要让batch膨胀带来的方差变化把你辛苦调好的训练曲线一夜打回原形。
返回列表