ARTICLE DETAIL

资讯详情

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

梯度累积实战:显存不足时扩大等效Batch Size的完整指南

梯度累积实战:显存不足时扩大等效Batch Size的完整指南 做深度学习训练的人几乎都会撞上同一堵墙显存不够。模型越来越大batch size越调越小最后小到loss曲线抖得没法看。梯度累积就是在这样的现实里被反复使用的一个技巧——它不额外占用显存却能让等效batch size翻几倍让训练在资源受限时依然稳定收敛。这篇文章要把我这些年用它踩过的坑、调参的经验、还有代码层面的细节一次性讲清楚无论你是刚跑通MNIST的新手还是正在跟大模型搏斗的工程师都可以拿来当参考。1. 为什么需要梯度累积显存不够时的第一选择1.1 显存瓶颈与batch size的现实矛盾先说说我最早遇到的场景。那时候我在调一个BERT微调任务显卡是单张2080Ti11G显存。模型塞进去一个合适的batch根本放不下。我第一反应是调小batch size从32一路降到8显存是够了但loss曲线开始剧烈震荡验证集表现也明显变差。这个现象很多人都有体会。batch size太小每个batch算出来的梯度噪声大参数更新方向不稳定训练像是喝醉了走路左一步右一步。特别是对Transformer这类模型batch size一降统计估计和训练稳定性都会受影响。可硬件限制就摆在那总不能为了训个模型去换一张A100。梯度累积解决的就是这个矛盾。思路非常朴素既然显存里放不下大batch那就分几次前向反向把梯度先“攒着”攒够了再更新一次参数。从数学上看把一个batch分成N份先后处理累积N份梯度后更新其梯度均值和一次性用N倍大的batch更新参数是一致的。这个场景是梯度累积最常见的需求硬件受限、需要大batch、又不想牺牲训练稳定性。还有一种情况是做对比学习batch越大正负样本的对比信息越丰富但显存撑不住这时候梯度累积几乎是唯一不动模型结构就能扩batch的解法。1.2 梯度累积到底在做什么说句实在话梯度累积这个名字听起来挺唬人但底层逻辑就是一段简单的循环控制。正常训练时每个batch都完整走一遍“前向、反向、更新参数”三步梯度累积时前向和反向照常做但参数不急着更新而是等步数达到设定值N之后才统一踩一次optimizer.step()然后清零梯度重新开始。这里有个关键前提很多人没意识到PyTorch里梯度默认是“累加”的。每次loss.backward()计算出的梯度会加到param.grad上而不是覆盖。平时训练每个step先zero_grad()就是为了把这个累加效果清掉。梯度累积恰恰是利用了这个机制——在N个小批次之间不清理梯度让它们自然叠加最后一步才更新。我习惯用一个生活化的比喻帮助理解普通训练像是每凑够一锅菜就炒一锅而梯度累积是把N锅菜的食材先都洗净切好全部堆在一起最后一次性下锅。食材总量没变但炒出来的菜味道更稳定。对应的就是等效batch变大了梯度估计更准了。2. 原理拆解从反向传播到参数更新的完整链路2.1 梯度从哪来、到哪里去要真正用好梯度累积不能只停留在“会抄代码”的层面得搞清楚梯度在内存里的流动过程。一次loss.backward()执行完毕后模型里每个参数param都会在.grad属性中存下一份梯度。这份梯度是该参数对当前batch损失函数的偏导数。假设模型参数为θ一个小批次的损失为L那么反向传播得到的就是∂L/∂θ。正常情况下紧接着用optimizer.step()执行θ ← θ − η · ∂L/∂θ然后zero_grad()把梯度归零。而梯度累积模式下backward()之后什么都不做直接进入下一个小批次的前向。第二次的梯度∂L/∂θ又加到同一个.grad上变成∂L/∂θ ∂L/∂θ。如此往复N次.grad里存的就是N个小批次梯度的和最后optimizer.step()更新一次。这个过程里最忌讳的操作就是在每个小批次后面随手加一句zero_grad()。加了之后梯度就被清空累积效果不复存在参数更新的频率变成每个小batch一次——相当于你白忙活还多付了N-1次额外的优化器开销纯属给自己添堵。还有一种细节值得注意梯度累积并不会改变模型的前向计算路径也不会改变BatchNorm的运行方式。模型结构和loss的计算逻辑跟普通训练完全一致只是参数更新的时机被延后了。2.2 等效batch size是怎么算出来的梯度累积的核心价值是“等效batch size”计算公式很简单等效batch size 单次小batch size × 累积步数N举个例子你显存只能放下batch size为16的数据设累积步数N4那么等效batch就是64。如果你的训练还用了多卡数据并行公式再加一项等效batch size 单卡batch size × 卡数 × 累积步数N我自己常用的一个判断标准是先在不考虑显存限制的理想情况下确定一个“目标batch size”比如ResNet训练一般用256ViT这类模型可以用到1024甚至更大然后根据单卡能塞下的最大batch反推需要多少步累积。这里有一个容易想当然的误区等效batch size只是在“梯度均值”层面等效并非在全部层面等效。模型参数确实每N步才更新一次每一步使用的是N份梯度的和或均值但不同小batch之间的数据顺序、BatchNorm统计量的更新频率跟真正的大batch训练还是有细微差异。这些差异通常不影响大局但在某些场景会成为问题后面第3章会详细展开。2.3 为什么不直接调小batch size我经常被人问既然显存不够直接把batch调小不行吗何必搞梯度累积这么麻烦这个问题的答案藏在梯度噪声里。每个batch的梯度都是对全体数据真实梯度的有偏估计。batch越小估计方差越大。方差大了之后参数更新的轨迹就会比较曲折训练需要更小的学习率来补偿训练时间拉长甚至可能收敛到很差的局部最优点。深度学习训练里稳定性和收敛速度往往比单次迭代的绝对速度更重要。batch太小还会带来一个更隐蔽的问题很多论文里的超参数比如学习率、weight decay、warmup步数都是基于特定batch size调出来的。你贸然把batch减半原来的学习率大概率要跟着变整个训练配置全乱套了。梯度累积能让你在不改动原有超参体系的前提下把“等效batch”还原到论文推荐的水平这在实际复现工作里是非常重要的。不过也要说句公道话梯度累积并非银弹。如果显存连一个合理的小batch都放不下或者你的显存只够塞下batch size为2、3这样极端小的数据梯度累积的收益就大打折扣因为单次前向的统计意义太弱累积很多步也只是把一堆噪声加在一起。这种情况下更好的选择是换更小的模型、用梯度检查点、或者混合精度训练而不是一味依赖累积。3. 实操落地PyTorch实现梯度累积的完整方案3.1 最简实现几行代码完成累积逻辑先上一个最朴素、最容易理解的写法。假设你的数据加载器每个iteration产出一个batch希望在4步累积后再更新参数accumulation_steps 4 for step, (inputs, labels) in enumerate(train_loader): outputs model(inputs) loss criterion(outputs, labels) # 关键点不除以累积步数也可以但除以之后梯度量级更接近真实大batch loss loss / accumulation_steps loss.backward() # 每累积满 accumulation_steps 步才更新一次参数 if (step 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这段代码核心就两件事loss.backward()在累积if (step 1) % accumulation_steps 0在控制更新时机。等训练结束时如果最后几个batch没凑满N步梯度会残留在.grad里需要手动补一次optimizer.step()清理掉。第一次上手的人最容易犯错的地方是把zero_grad()放在了每个iteration都执行的位置。记住zero_grad()只应该出现在参数更新的那个分支里。还有一个常见错误是忘了在训练结束时处理残余梯度下次训练继续加载优化器时上一次残留的梯度会干扰第一个step的更新。稳妥的做法是每个epoch结束时都确认梯度已经清干净。3.2 loss归一化除以累积步数的说法与依据关于要不要把loss除以accumulation_steps江湖上有两派说法我把逻辑捋清楚。如果不除以N累积N步之后的梯度和就是N份梯度的和相当于梯度放大了N倍在优化器步长固定的前提下实际更新幅度等效于学习率乘以N。这在某些场景下能无意识加速训练但很容易导致初期loss爆炸尤其是用Adam这类带自适应学习率的优化器时梯度量级的突变会让优化器的状态估计紊乱。如果除以N累积N步后的梯度等于N份梯度的平均和真正batch size扩大N倍的梯度平均在数值上是一致的。这是更符合“等效batch size”语义的做法也是我个人的默认选择。代码上只需多写一行loss criterion(outputs, labels) / accumulation_steps loss.backward()这里要注意的是除以N之后单个小batch的loss数值本身会变小如果loss曲线里还同时使用了warmup等策略你看到的loss曲线会显得比普通训练小一些但这不是bug只是参考系不同。想要对比效果时记住把累积后的loss乘回N或者直接观察验证集指标。3.3 混合精度下的注意事项如果你同时使用混合精度训练AMP梯度累积有几个坑必须注意。PyTorch的torch.cuda.amp.GradScaler会对loss做缩放并在反向传播后自动处理梯度缩放。配合梯度累积时常见的写法是这样scaler torch.cuda.amp.GradScaler() for step, (inputs, labels) in enumerate(train_loader): with torch.cuda.amp.autocast(): outputs model(inputs) loss criterion(outputs, labels) / accumulation_steps scaler.scale(loss).backward() if (step 1) % accumulation_steps 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad()这里有个容易出问题的地方scaler.scale(loss)会对缩放后的loss再乘一个scale因子反向传播得到的梯度也是缩放过的。当你累积了N步缩放过的梯度后最后scaler.step()里会自动除以scale这没问题但问题是如果N步之间scale因子发生了更新当出现inf/nan时scaler会自动调整scale前后梯度可能不在同一个scale体系下。实际操作中我的建议是把GradScaler的更新也放在累积分支里尽量不要让scale在累积过程中变动。如果你的模型经常出现梯度溢出与其指望scaler救场不如先确认loss归一化是否正确、batch数据里是否有异常值。提示混合精度下的梯度累积最稳妥的验证方式是在累积分支结束后打印一步的梯度范数看看是否跟同配置的普通训练接近。差异过大就先别上模型训练先把这个数字对齐了再说。4. 梯度累积与学习率关键参数怎么调4.1 学习率缩放的经验法则梯度累积改变了“一次更新对应的数据量”所以学习率的设置逻辑也要跟着变。最经典的参考是线性缩放法则batch size翻K倍学习率大致也翻K倍。如果原来用batch size64、学习率1e-4训练很稳定现在用batch size16配合累积步数4等效batch还是64那么学习率应该保持在1e-4附近如果你刻意把等效batch从64提升到128那么学习率可以尝试1.5e-4到2e-4。这里面的数学直觉是更新次数变少了因为每N步才更新一次为了在总资源相同的情况下走完整段训练每一步应该迈得更大一些。但步子迈太大又会不稳定所以线性缩放法则有一个适用范围一般建议等效batch翻倍时学习率最多翻倍不要超过这个比例。我在实践中发现学习率缩放是否合适可以从训练初期的loss下降速度来判断。如果学习了前几百步loss就震荡得厉害说明学习率偏大如果loss下降明显变慢那可能是学习率偏小。1200步左右的观察窗口已经足以做出大致判断。4.2 不同任务场景下的策略选择不同模型的“学习率敏感度”差别很大我按经验分几类说明。对CNN类模型ResNet、EfficientNet等学习率相对好调线性缩放法则基本适用累积步数在2到8之间loss曲线通常不会出大问题。Transformer类模型要小心得多这类模型对学习率的敏感度很高Adam、AdamW搭配warmup几乎是标配。增大等效batch之后warmup步数也应该按比例增加否则前期学习率爬升太快容易导致loss飞掉。对比学习、自监督训练是另一个典型场景。这类方法天然追求大batch比如SimCLR在原始论文里用到了8192的batch size普通实验室根本达不到。这时梯度累积是无奈的但又是有效的替代方案。不过要注意对比学习的loss函数通常包含一个batch内样本相互区分的项batch size小的话负样本数量不足模型学到的表征质量会直接下降。我建议这种情况把可用的集显存先堆满batch再算需要累积多少步达到目标而不是漫无目的地设一个很大的N。剩下的经验性建议就一句话改动等效batch之后先只动学习率和warmup其他超参按兵不动跑一个完整epoch再看趋势。一次别改三个超参不然出了问题你根本不知道是谁的锅。4.3 学习率调度器的正确接法学习率调度器scheduler在梯度累积场景里是一个很容易踩雷的细节。PyTorch的torch.optim.lr_scheduler通常在每个optimizer.step()之后或之前调用lr_scheduler.step()。如果你的代码在梯度累积分支外错误地每个iteration都调用lr_scheduler.step()那么学习率衰减的速度会比正常训练快N倍训练后期几乎学不动。正确的做法是把调度器的step()也放在参数更新分支内和optimizer.step()放在一起if (step 1) % accumulation_steps 0: optimizer.step() scheduler.step() optimizer.zero_grad()从这里还能引出一个更底层的认知梯度累积并不会改变“参数更新总量”这一物理量。普通训练10000步更新10000次累积步数4后同样跑10000个batch更新次数变成了2500次。总的训练数据吞吐没变loss曲线的横坐标含义发生了变化。如果你的代码里到处用的是“iteration数”而不是“epoch数”那么开启梯度累积后所有跟步数相关的逻辑——warmup长度、调度器衰减周期、log打印频率——都要重新审视。我最常看到的问题就是有人开了梯度累积后忘了调整warmup步数结果模型在前几个epoch就出现了明显的不收敛。排查一圈才发现是调度器在每个iteration都衰减了一次学习率早早就降到了正常值的四分之一。5. 常见问题与排查技巧实录5.1 BatchNorm统计量的“表里不一”先说这个最隐蔽的问题。BatchNorm在训练时会维护一组running_mean和running_var它们是在每个batch前向过程中实时更新的与参数更新的时机没有直接关系。也就是说即使你做了4步梯度累积BatchNorm的统计量依然是每处理一个小batch就更新一次。当你把小batch size从64降到16、配合累积步数4时BatchNorm看到的仍然只是16个样本的统计量这会带来两个后果一是统计量噪声变大二是与“等效大batch”语义不一致——参数等效于大batch更新了但BatchNorm的估计还是小batch水平。这个问题在视觉任务里比较明显症状是训练loss正常下降但验证集泛化性能差或者训练后期出现震荡。解决思路有三个一是尽量保证单个小batch不要太小至少16个样本以上二是在网络支持的情况下把BatchNorm替换成GroupNorm、LayerNorm这类与batch无关的归一化方式三是如果你的训练脚本本身已经用了多卡同步BNSyncBN梯度累积下这个问题会被进一步放大务必额外注意。说实话我见过不少团队把“训练不稳”归咎于学习率折腾半天才发现问题是BatchNorm统计量跟等效batch不匹配。排查看loss曲线的同时建议把每个epoch的running_mean、running_var打出来看一眼如果波动远大于正常水平多半就是它的锅。5.2 梯度裁剪与梯度溢出梯度裁剪gradient clipping是Transformer训练的保命技能。但配合梯度累积时裁剪的位置和时机都有讲究。很多人习惯在每个iteration后面加一句torch.nn.utils.clip_grad_norm_在梯度累积模式下如果N个micro-batch的梯度还没垒齐就裁剪效果会变得非常诡异前几步梯度被压缩到很小的量级最后一步完整梯度又暴涨训练非常不稳定。正确的做法是只在optimizer.step()之前裁剪一次也就是作用于累积完的完整梯度if (step 1) % accumulation_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() optimizer.zero_grad()另外如果你用了loss归一化除以N那么梯度范数大致落在了正常量级max_norm可以沿用普通训练时的经验值如果没有归一化梯度范数会放大N倍max_norm需要相应调大否则裁剪会过度干预训练。这个细节我建议直接记住用除法归一化的人裁剪参数不用改不归一化的人先看看梯度范数再决定。5.3 DDP下的同步开销问题上了多卡分布式训练DDP梯度累积还有个额外的性能陷阱。DDP在反向传播时会通过AllReduce同步各个卡上的梯度。默认情况下每个loss.backward()都会触发一次梯度同步。如果你累积4步就意味着4次反向传播做了4次不必要的AllReduce而通常你只希望最后一步才同步。正确的做法是使用model.no_sync()上下文管理器在前N-1个micro-batch反向传播时跳过梯度同步最后一个batch再正常反向并且触发同步accumulation_steps 4 for step in range(len(loader)): inputs, labels next(iter(train_loader)) for i in range(accumulation_steps - 1): with model.no_sync(): loss criterion(model(inputs), labels) / accumulation_steps loss.backward() # 最后一步正常反向触发AllReduce loss criterion(model(inputs), labels) / accumulation_steps loss.backward() optimizer.step() optimizer.zero_grad()这个写法省掉的通信量非常可观。假设4卡、4步累积优化的通信次数直接从4次降到1次训练吞吐可能提升10%到20%。我一度以为自己的多卡训练速度慢是网络带宽问题后来才发现瓶颈在这里。5.4 常见问题速查表现象大概率原因处理方式loss震荡剧烈、不收敛学习率偏大或未按等效batch调整按4.1的线性缩放法则重新设定loss在前期就飞掉warmup/调度器步数未调整将调度器和warmup按更新步数重算每个iteration显存逐步上涨至OOM梯度累积过程中zero_grad()位置错误检查zero_grad()是否放进了更新分支训练loss正常但验证效果差BatchNorm统计量与等效batch不匹配增加单卡batch、更换归一化方式多卡训练吞吐异常低DDP下每次backward都触发AllReduce使用model.no_sync()跳过前N-1次同步梯度范数异常大或出现NaNloss未归一化且学习率偏大归一化或调大max_norm并降低学习率训练到最后到底没收敛调度器衰减过快确认scheduler.step()只在更新分支调用6. 进阶玩法与个人心得6.1 动态梯度累积的探索固定步数的梯度累积只是入门实际操作里还可以做动态调整。一个常见的思路是训练初期梯度噪声大需要更大的等效batch来稳住方向这时候累积步数可以设大一些训练后期模型接近收敛更新可以更频繁于是逐渐减小累积步数。实现上并不复杂用一个变量控制累积步数按epoch或按数据量动态更新即可def get_accumulation_steps(epoch): if epoch 10: return 8 elif epoch 30: return 4 else: return 2我试过几次动态累积在有明显阶段划分的任务上能带来一两点的精度提升但也让训练流程复杂了不少日志、断点恢复、调度器全都得跟着适配。个人建议除非你确实面对精度瓶颈否则先老老实实用固定步数跑通全流程动态策略属于“锦上添花”的优化而不是“雪中送炭”的必需方案。6.2 需要避开的几个误区最后把梯度累积最常见的错误认知集中列一遍都是我在各种代码评审里反复遇到过的。第一个误区是以为梯度累积能无限增大等效batch。实际上当累积步数超过8到16之后收益会急剧递减。因为梯度是数据分布的估计估计的方差下降速度与样本量是平方根关系——batch从64到128提升明显从1024到2048就没那么明显了再往上就纯粹是浪费训练时间。第二个误区是忽略数据顺序对累积效果的影响。梯度累积要求N个micro-batch尽量独立均匀。如果你的数据加载器在每个epoch开始时把数据按标签排序前几步累积的可能是同一类样本累积出来的梯度会有偏。解决办法是启用shuffle这虽然是最基础的训练常识但偏偏在梯度累积场景里最容易被人忽略。第三个误区是拿“显存不够”当借口把所有问题都甩给梯度累积。有些时候小batch就不应该硬凑成大batch换成更小的模型、用梯度检查点、减小输入分辨率效果可能还更好。梯度累积是工具箱里的一件常用工具把它当成万能钥匙迟早会在某个任务上翻车。6.3 关于断点续训与复现的提醒既然你已经在用梯度累积相信很快也会遇到断点续训的问题。保存checkpoint的时候优化器的状态、学习率调度器的状态、当前step位置都要一并保存。梯度累积本身没有额外状态但有一个很容易丢的细节如果程序在累积进行到一半时中断你恢复训练时需要考虑是否要先丢弃当前残留的梯度。最稳妥的做法是checkpoint里记录当前epoch和step恢复时先把optimizer.zero_grad()清一次重新从完整的一轮累积周期开始。别看这是个小细节我见过有人因为这个残留梯度的问题恢复了两次训练但结果一次比一次差查了半天才发现是累积计数错位。最后再分享一个我的习惯凡是开启梯度累积的训练脚本我一定会在启动后打印前两步的梯度范数和optimizer.param_groups[0][lr]。这两个数字正常后面出问题的概率就小一半。训练这种事多一道保险总比事后复盘强。
返回列表