ARTICLE DETAIL

资讯详情

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

梯度累积原理与实战:突破显存限制的大模型训练技巧

梯度累积原理与实战:突破显存限制的大模型训练技巧 我最近在调一个7B模型单卡batch size最多只能塞下2个样本想用32的batch根本不可能。换成梯度累积之后per_device_batch_size2、gradient_accumulation_steps16等效batch到了32训练才真正跑起来。梯度累积这个技巧几乎所有做过大模型微调、跑过CV分割或者受限显存训练的人都会遇到但很多人只是抄了代码并不知道它到底在做什么、为什么这么做、学习率要怎么配合。这篇文章就把梯度累积从头到尾讲清楚从原理到代码再到最常见的坑一次说透。1. 梯度累积到底解决什么问题1.1 核心场景batch size上不去的困境先聊聊为什么需要这个技术。深度学习训练中batch size几乎是最重要的超参数之一它直接影响梯度估计的稳定性、BatchNorm等层的统计量、以及最终的收敛效果。但实际训练时batch size不是你想设多大就设多大限制主要有几个显存限制。显存里不只有数据和模型参数还有每一层的激活值、临时buffer以及优化器状态。拿Adam来说每个参数要额外存一阶动量、二阶动量再加上梯度本身一套下来显存开销是模型参数体积的几倍。模型越大能塞下的batch就越小。数据加载与预处理内存限制。有的场景比如医学影像、视频理解单个样本就已经很大一个小batch就会吃满显存。分布式同步成本。多卡训练时大batch对应更大的梯度AllReduce通信量batch太大会明显拖慢单步速度。在这些限制下如果你坚持用大batch要么等着OOM要么只能缩小batch但缩小batch之后梯度噪声变大lr要跟着调训练稳定性变差。梯度累积就是在这种情况下被广泛使用的折中方案。1.2 梯度累积做了什么梯度累积的思想特别直白把一次大batch的更新拆成多个小batch的多次前向反向每次不更新参数而是把算出来的梯度累加到一起等累加了足够多次之后再统一做一次优化器更新。举个例子你想用batch size32训练但显存只允许batch size4。那么可以这样操作取4个样本前向计算loss反向传播得到梯度。先不更新参数把梯度加到参数.grad里面。再取下一个4个样本继续前向反向继续累加梯度。重复8次。8次之后把累积的梯度除以8或者用其他方式归一化然后执行一次optimizer.step()再清空梯度。这就是所谓的“用gradient_accumulation_steps8模拟batch size32”的基本过程。注意我在这里特别写了“除以8”这个细节很多教程都会漏掉这一点后面会详细展开。1.3 一个容易被误解的点这不是简单地把更新推迟有人可能会想不就是攒几轮梯度再更新一次吗我把optimizer.step()的频率降低不就行了实际上没有这么简单。如果只是推迟更新那你每次反向传播的梯度是各自micro batch算出来的它们的量级和一次真实大batch的梯度并不一样。用4个样本算出来的梯度数值一般来说比用32个样本算出来的梯度大或者小取决于怎么归一化中间还涉及BN的running stats怎么更新、梯度裁剪用在哪一步、学习率调度怎么计算等等。所以梯度累积并不是“少调几次优化器”而是一种需要精心管理的梯度聚合过程。我用一个搬运的类比来理解你要把一吨货从A搬到B一次运不完就拆成20趟每趟50公斤。但你不是每搬一趟就发车而是先把货堆到卡车上堆到一吨再发车。这个类比能解释“分步搬运、集中运送”的直觉但注意梯度不是货物那么简单每趟搬的时候来自不同数据子集数据的分布会有差异而且你还需要决定是堆满一吨发车求均值还是直接发一辆超载车求和这就是后面要讲的数学细节。2. 原理拆解为什么累积梯度不等于直接加大batch2.1 数学上到底发生了什么先假设一个标准训练设置。一个batch里有N个样本标准的batch loss是L (1/N) Σᵢ ℓ(xᵢ, yᵢ)对应的梯度是对L求导也就是对每个样本的梯度求平均g (1/N) Σᵢ ∇ℓ(xᵢ)现在把N个样本拆成k组每组m个样本Nk*m。第j个micro batch的梯度是hⱼ (1/m) Σᵢ∈Bⱼ ∇ℓ(xᵢ)注意这里的hⱼ已经是“平均梯度”不是求和梯度。如果你只是把k个hⱼ直接累加得到的是h_sum Σⱼ hⱼ (1/m) Σᵢ∈All ∇ℓ(xᵢ) k * g也就是说简单的累加相当于把梯度放大了k倍。k个micro batch的梯度总和与标准batch梯度之间差了一个k的倍数。如果希望严格等价于真实大batch的平均梯度应该做的是h_avg (1/k) Σⱼ hⱼ (1/(k*m)) Σᵢ ∇ℓ(xᵢ) g这就解释了为什么很多实现里会看到loss loss / gradient_accumulation_steps这样的代码。除以这个k正是要把“多个micro batch平均梯度求和”再拉回“全局平均梯度”的水平。很多人抄了代码但不知道为什么除以accumulation steps就是这里的问题。2.2 “均值累积”与“求和累积”只差一个除法结果差很多基于上面的推导代码实现上其实有两种流派均值累积mean accumulation每个micro batch的loss先除以k再反向梯度累加后天然等于标准大batch梯度。这也是HuggingFace Trainer、PyTorch官方例子里最常见的方式。求和累积sum accumulation不除以k直接把k个micro batch的梯度相加。这种情况下最终梯度约等于标准大batch梯度的k倍实际上相当于把学习率放大了k倍或者说把loss函数定义成了“batch内平均损失求和累加”的形式。很多人写手动实现时用的是第二种最终训练也能跑因为优化器大概率会把梯度scale吃掉一部分。但问题在于SGD、momentum、Adam对梯度scale的敏感程度完全不同你在求和模式下换一个优化器训练表现会非常不一样而均值模式的行为则更可预测、更接近直接用大batch的效果。我的建议是除非你有明确理由比如某些论文明确写了不归一化否则优先用均值累积也就是在loss.backward()之前把loss除以accumulation steps。这样调参的时候你可以直接参考“等效batch size”的经验不用心算梯度被放大多少倍。2.3 学习率如何缩放学习率与“等效batch size”的关系是很多人最容易翻车的地方。搜索引擎里“梯度累积学习率”是高频搜索词说明大家都被这个问题卡过。如果你的模型原来用batch sizeB、学习率lr现在用梯度累积把等效batch size变成B B * k学习率要不要跟着调理论上通常有两种规则线性缩放规则linear scaling rulelr_new lr * (B / B) lr * k。直观解释是batch变大后梯度估计更稳定噪声更小所以可以走更大的步。这个规则在SGD/momentum下表现很好。平方根缩放规则sqrt scaling rulelr_new lr * sqrt(B / B) lr * sqrt(k)。相对更保守在Adam类优化器下经常比线性缩放更稳。实际经验是什么我自己在AdamW下做实验线性缩放设置lr*4的情况loss曲线经常直接起飞尤其是模型较大、warmup不够长的时候。平方根规则通常更稳但也不是万能——如果你的等效batch是通过gradient_accumulation_steps很大比如32甚至64实现的我建议先保持原lr跑100步观察梯度norm和loss变化再逐步提高lr而不是一上来就按公式硬套。这里还有一个关键点学习率调了之后学习率调度器也要跟着改。如果你原来用的是一个step-wise的训练计划总步数数据量/单batch batch_size引入梯度累积之后真正的optimizer.step()次数只有原来的1/kwarmup步数、cosine周期都要按新的总步数重新算。很多人只改了lr没改schedule结果预热阶段还没结束cosine已经衰减到一半了训练自然对不上节奏。2.4 梯度累积与学习率的关系总结整理成一张速查表场景推荐做法原因原来SGD训练等效batch翻k倍lr乘k梯度更稳可大步更新原来AdamW训练等效batch翻k倍lr乘sqrt(k)或不动观察后再调Adam自适应学习率对scale不敏感激进调lr容易炸warmup步数按optimizer.step()总数重新计算累积后总更新次数减少预热比例要重新对齐cosine底部按optimizer.step()总数重新计算否则衰减节奏错位小模型小数据集保守调小lr过拟合风险高累积带来的梯度稳定增益有限3. 代码落地手写实现和框架内置3.1 最朴素的PyTorch手写版本先给一个最基础、也最推荐照抄的版本accum_steps 8 optimizer.zero_grad() for step, batch in enumerate(dataloader): loss model(batch) / accum_steps # 关键先除以累加步数 loss.backward() if (step 1) % accum_steps 0: # 可选在这里添加梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() optimizer.zero_grad()这段代码里三个细节值得说明loss model(batch) / accum_steps这就是前面讲的均值累积。如果不写这一步你的梯度会放大accum_steps倍。optimizer.zero_grad()和optimizer.step()每一轮只执行一次都在if判断里面。如果忘了zero_grad或者把zero_grad放在了if外面梯度就会在多个accum cycle之间累积结果完全错误。梯度裁剪放在optimizer.step()之前、且在整个accumulation cycle完成之后。如果每个micro batch都做一次裁剪会让梯度方向被干扰而且裁的是局部梯度而不是最终要用的累积梯度效果完全不同。另外一个细节optimizer.zero_grad(set_to_noneTrue)比optimizer.zero_grad()更省显存因为前者直接把梯度置为None而不是分配一块零张量。在PyTorch 2.x里这个优化几乎无成本能写就写。3.2 HuggingFace Trainer直接用参数如果你用Transformers库不需要手写循环。直接在TrainingArguments里配置from transformers import TrainingArguments training_args TrainingArguments( output_dir./output, per_device_train_batch_size2, gradient_accumulation_steps16, learning_rate1e-5, warmup_ratio0.03, lr_scheduler_typecosine, fp16True, logging_steps1, )Trainer内部会在拿到每个micro batch的loss之后自动除以gradient_accumulation_steps再在满足步数条件时执行optimizer step。这也是很多人感觉“我明明什么都没干为什么行为正确”的原因——因为框架帮你做了均值累积。但注意框架并不会自动帮你调学习率。learning_rate1e-5这个值是我自己设定的框架不知道你之前用的是多少batch也不会因为accumulation steps变大而自动放大lr。所以用Trainer时最容易犯的错误是只改了gradient_accumulation_steps其他参数什么都不动然后就觉得训练变慢了或者变快了其实都在于lr没有配套调整。3.3 分布式训练与梯度累积不会互相冲突吗很多人一开始用DDP 梯度累积的时候会怀疑梯度是否被重复缩放。这里把逻辑理清楚。DDP训练时每张卡拿到的micro batch是本地数据前向反向之后梯度会在所有卡之间做一次AllReduce默认是平均。然后本地把不同卡的梯度平均之后用这个平均梯度更新本地模型的副本。这种情况下如果你又用了梯度累积一定要理解DDP的AllReduce平均发生在每个micro batch的backward时而不是在accumulation step之后。也就是说每个micro batch累积的梯度已经是“全局卡数平均之后”的梯度。多个micro batch累加之后再做一次均值累积除以accumulation steps得到的是“多轮micro batch的全局平均梯度之和再平均”这与真实大batch的更新方向一致。这里有一个性能层面的大坑如果不做任何特殊处理每个micro batch的backward都会触发一次AllReduce通信。你虽然把参数更新的频率降低了但通信频率没有降低训练吞吐可能并没有提升多少。正确做法是在非最后一个micro batch时用model.no_sync()上下文禁止梯度同步只在最后一个micro batch结束时正常同步from torch.nn.parallel import DistributedDataParallel as DDP model DDP(model) for step, batch in enumerate(dataloader): # 非最后的累积步跳过梯度同步 context model.no_sync() if (step 1) % accum_steps ! 0 else nullcontext() with context: loss model(batch) / accum_steps loss.backward() if (step 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()这段逻辑在Transformers Accelerate环境下也有对应实现但手写时特别容易被忽略。如果你发现梯度累积后训练速度并没有比正常大batch快多少先检查是不是每个micro batch都在做AllReduce。3.4 AMP混合精度和梯度裁剪的配合梯度累积与混合精度一起用的时候有个容易踩的细节GradScaler。AMP训练时梯度会通过loss scaling放大若干倍避免FP16下梯度下溢。如果你在每个micro batch都调用scaler.step(optimizer)但实际又没到累积步数优化器状态根本没更新scaler的scale因子会反复调整非常不稳定。正确做法是只在完成一个完整accumulation cycle时调用scaler.step()和scaler.update()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_norm1.0) scaler.step(optimizer) # 真正的更新点 scaler.update() optimizer.zero_grad()要注意scaler.unscale_(optimizer)的顺序。如果你想做梯度裁剪并且用了AMP最好先调用unscale_把梯度还原再做clip否则clip的阈值是在缩放后的尺度上计算的你以为是1.0实际作用在128倍缩放过的梯度上几乎等于没有裁剪。4. 梯度累积的效果怎么评估4.1 怎么判断你写对了很多人写完梯度累积代码训练loss也在降但不确定自己写得到底对不对。我建议用一个小实验快速验证。取一个很小的模型、固定随机种子分别跑两种配置配置Abatch size4不用梯度累积。配置Bbatch size1gradient_accumulation_steps4。理论上如果用均值累积两个配置的梯度方向应该非常接近loss数值不一定完全一样因为数据顺序如果不同梯度会有差异使用相同顺序、相同shuffle seed时可以逼近完全一致。你在配置B里打印每一个micro batch的loss再打印最终更新时的梯度norm对比配置A的梯度norm。如果配置B的梯度norm大约等于配置A的梯度norm说明累积逻辑正确。如果大了一倍或四倍说明你没有做均值累积。对数梯度norm还有一个额外好处你可以观察norm的绝对值来判断lr是否偏了。梯度norm长时间在一个数量级上不下降说明模型可能根本没在学突然飙升几个数量级往往是梯度爆炸或lr过大。4.2 与真实大batch仍不完全一样虽然梯度累积在更新公式上可以等价于真实大batch但在实际训练中存在三类差异需要心里有数。第一类是BN等统计量。BatchNorm层的running mean和running variance是根据每个micro batch的数据统计出来的不是根据“等效的完整大batch”统计出来的。你用batch size1累积16步和直接batch size16跑一次BN的统计量不会完全相同。这也是为什么梯度累积在CV任务上的“等效性”比在NLP任务上更差一点尤其是图像分辨率大、batch size很小的时候。如果你的模型里有BN且训练和测试batch差异较大建议谨慎依赖梯度累积模拟大batch。第二类是数据顺序与梯度平滑。真实大batch一次看到32个样本梯度是这32个样本的平均方向梯度累积是分16次看到32个样本中间插入了backward虽然最终梯度近似但网络内部的dropout、数据增强、每个micro batch的随机顺序会产生额外的噪声。这些噪声对最后几轮的收敛会有影响部分实验会发现梯度累积的最终效果比真实大batch略差一点也可能在某些场景下略好——因为它天然带有一定正则化效果。第三类是优化器行为。Adam类优化器在每个参数上维护一阶和二阶动量估计。真实大batch一次更新用的是该批的梯度梯度累积多次更新时虽然最终的是一个等效梯度但动量的更新时序和方差特性不一样导致最终的参数轨迹不会完全一致。这也就是为什么lr缩放规则在Adam下并没有一个普适公式。4.3 哪些任务最受益、哪些要谨慎从我的使用经验来看梯度累积在下面这些场景收益最大大模型LLM微调。显存是硬约束batch size动不了梯度累积几乎是标配。高分辨率图像分割/物体检测。一张图可能就占几个G显存不累积根本跑不了。小batch导致BN统计不稳定、loss抖动的场景。通过累积把等效batch提上去能明显平滑loss曲线。需要谨慎的场景也有小数据集上训练时间本来就很短。引入大accumulation step会让参数更新次数很少模型可能还没学好就结束了。超长序列训练。比如LLM上下文长度推到8192以上单个样本已经很大如果累积步数也设得很大训练一步要跑很久出问题排查效率很低。模型内部有依赖全局统计信息的层比如某些归一化层、对比学习里的超大负样本梯度累积模拟不了这种“全局batch”的意义。5. 常见问题速查5.1 一张表解决大部分问题现象可能原因处理建议训练loss爆炸/NANlr过大或没有做均值累积先看梯度norm关掉lr缩放把lr降到原值或更低训练loss下降极慢梯度没被正确累积或lr过小打印每个accum cycle结束时的梯度norm确认是否有数据显存仍然不够优化器状态太大或activations太大用gradient checkpointing把per_device_batch_size降到1多卡训练吞吐没提升每个micro batch都在做AllReduce加上model.no_sync()混合精度下loss异常GradScaler在micro step被反复update确保只在完整accumulation后调用scaler.step/update每个epoch最后一个batch没有更新数据量不能被accum_steps整除直接在余数上执行一次step或在dataloader里drop_last用Trainer改accum后训练变慢有效step次数减少lr schedule不匹配重新设计warmup和lr schedule检查梯度同步是否过于频繁5.2 几个值得单列的排查经验先说说数据量不能被accum_steps整除的问题。假设数据集有3000条per_device_batch_size4accum_steps8一共75个真正的optimizer step。这么算完全没有问题。但如果你用的数据加载器自动drop掉不足一个accumulation cycle的余数很可能每轮训练就少掉几十个样本。可以接受但要注意如果你同时用distributed sampler多卡上每个卡的数据量本身已经不均drop之后可能出现“某些卡比别的卡多跑一个cycle”的微妙不同。手写训练循环时我一般会把dataloader设置为drop_lastTrue并确保总数据能整除accum_steps省心很多。再说一个很隐蔽的坑accum_steps过大的时候checkpoint中的global_step含义会变。你在Trainer里设了logging_steps1你会看到log的step不是按micro batch来的而是按有效的optimizer step来的。有些人在恢复checkpoint继续训练时发现learning rate不对就是因为在恢复时把step数理解错了把micro-batch step当成了optimizer step。检查一下Trainer保存的trainer_state.json里的global_step对照一下自己的预期。还有关于梯度延迟。在动手调参之前我强烈建议先在训练脚本里加上一个梯度norm的hook或日志total_norm 0.0 for p in model.parameters(): if p.grad is not None: total_norm p.grad.norm().item() ** 2 total_norm total_norm ** 0.5 print(fstep {global_step}, grad norm: {total_norm:.4f})梯度norm是判断训练是否健康的最直接信号。它长时间不变化说明优化器可能没吃进去梯度它突然暴涨说明lr或者梯度累积逻辑出了问题它归于零说明你可能已经收敛或者梯度消失。这个日志能帮你省下大量调试时间。6. 一点个人建议我自己的调参习惯是能用真实batch size尽量用真实batch不要为了梯度累积而梯度累积。真到显存不够时优先把per_device_batch_size降到1再用accum_steps补。补多少合适我通常在LLM微调时用accum_steps8到32之间lr会在原基础上先乘sqrt(accum_steps)试跑200步看loss曲线和梯度norm再决定要不要加大。如果中途出现发散第一件事不是降lr而是检查梯度累积有没有正确除以步数——这个错误我在早期手写代码时踩过不止一次每次都是loss直接冲上Inf才发现。梯度累积本身不复杂把它背后的数学和工程细节理清楚之后你会发现在调度lr、设置warmup、判断训练是否健康时都更有把握了。
返回列表