news 2026/10/2 2:50:43

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
梯度累积:大模型训练显存不够时的关键技巧与实战指南

跑过大模型训练的人,大概率碰到过这样的场面:一开训练就OOM,被迫把batch_size从32砍到8,loss曲线抖得像心电图。这种时候,老手通常都会说:试试梯度累积(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曲线方差大、收敛慢,甚至优化器在局部震荡出不来。另一个常见副作用是BatchNorm:BN的均值和方差是“当前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_{i=1}^{B} \nabla \ell_i ]

然后更新参数:

[ \theta_{t+1} = \theta_t - \eta g ]

梯度累积的本质,是把这B个样本拆成K份,每份m个样本,先算每份的梯度 \(g_k = \frac{1}{m}\sum_{i=1}^{m} \nabla \ell_{k,i}\),再把K个梯度做平均:

[ G = \frac{1}{K}\sum_{k=1}^{K} g_k = \frac{1}{K \cdot m}\sum_{i=1}^{K \cdot m} \nabla \ell_i = \frac{1}{B}\sum_{i=1}^{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倍调整前提是配足够长的warmup
Adam/AdamW不调,或最多乘以sqrt(K)自适应机制已抵消部分方差变化,常用做法是保持lr不变
LoRA微调保持lr不变有效batch翻倍对PEFT影响较小,动lr反而容易破坏原本调好的基座
带BN的CNN + SGD按K倍线性调,但必须同步加大batch观察BN统计量只靠累积不一定能得到大batch的BN效果

举一个具体例子:我微调LLaMA类模型时,单卡batch=2,accum_steps=8,有效batch=16,学习率从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 batch,DDP就会同步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_batch=1、accum=8、8张卡时,一次参数更新实际吃掉的就是64个样本的梯度信息。这也是多卡训练里非常常见的配置。

4.3 多卡边界的三个隐藏坑

第一个坑是整除。如果你的数据总量不能被卡数乘accum_steps整除,最后几个不足K步的micro batch会搞乱循环:某几张卡凑不足K次backward,导致DDP在步数边界上hang住。稳妥做法是drop_last,把尾部长尾直接扔了,或者在末尾手动判断。

第二个坑是学习率视角。多卡下累积步数和卡数都放大有效batch,学习率要不要调,取决于你原本“有效batch”设计成多少,而不是只看累积步数。同样是accum=8,8卡和单卡的学习率策略很可能不同。

第三个坑是通信频率与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。如果你在累积中途频繁update,scale值会被不完整的累积梯度带偏,轻则训练不稳,重则连续出现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就clip,clip的对象还是“被放大过”的梯度,最终步长会偏差。这句是很多人查半天都没发现的隐藏点。

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时,shuffle=True的逻辑是每个epoch开始前打乱一遍全局顺序,循环内取到的batch本身是连续的,所以梯度累积代码里并不会自动重复抽样。但如果你自己写数据流,或者用那种“每个step重新随机抽”的数据接口,就得注意:累积的K个micro batch必须来自同一批采样策略,且互不重叠,才能保持等价性。

5.4 BatchNorm的统计量不会因为累积而变大

这是梯度累积最容易被误解的一点。很多人以为把batch=8累积8步,就等价于batch=64训练。数学上梯度确实接近,但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 batch=1、accum=16到64,让显存塞满训练、时间换有效batch。

6.2 不该硬用梯度累积的情况

如果显存充裕,直接大batch训练永远是更稳的选择。梯度累积虽然梯度等价,但参数更新频率降低、BN统计量偏差、dropout随机性增加,这些都会让最终效果和真正的大batch有细微差异。

此外,在线学习和流式场景不适合累积:每条数据都要即时更新模型,延迟步进反而有害。还有一些对BN特别敏感的任务,累积带来的收益有限,最优先方案是同步BN而不是梯度累积。

6.3 我的一点个人心得

这些年调模型,我对梯度累积的经验可以总结成三句话:先算清楚有效batch size有多少,再决定要不要动学习率;多卡必用no_sync,单卡也要记得把gradient clipping放到最后一个micro batch之后;AMP和累积代码尽量复用已经验证过的模板,不要随手改。

前阵子调一个7B模型的LoRA,单卡micro batch=1,accum=16,有效batch=16,学习率用2e-4没动,warmup拉到总步数的5%,整体曲线平滑得让我意外。反倒是另一组实验把有效batch翻到128,学习率同步调大了50%,结果几个step后loss直接飘了。重新把lr降回原始值,才恢复稳定。所以,我现在的默认操作是:累积步数越大,学习率调整越克制,宁可慢慢来,也不要让batch膨胀带来的方差变化把你辛苦调好的训练曲线一夜打回原形。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/2 2:50:43

AI论文写作Skills完全指南:从润色到降重的一站式技能包

最近在学术群里经常看到有人问:为什么别人用 AI 写论文,能直接给出可用的文献综述框架、能按期刊口味改句式,自己却每次都要从“帮我润色这段话”开始?其实差别往往不在模型,而在一个叫 Skills 的东西。如果你最近也在…

作者头像 李华
网站建设 2026/10/2 2:50:06

Linux正则表达式实战:从BRE/ERE到grep/sed/awk三剑客

Linux 正则表达式,说实话是很多刚接触命令行的人的第一道坎。我见过太多同事在 grep、sed 里被反斜杠和竖线绕得晕头转向,最后干脆放弃,回到图形界面里人工翻日志。这篇文章就是写给所有想在 Linux 命令行里高效处理文本的人——不管你是运维…

作者头像 李华
网站建设 2026/10/2 2:49:30

单节点Hadoop伪分布式集群搭建指南:从零配置到ZooKeeper整合

1. 项目概述与整体设计思路1.1 单节点集群到底在解决什么问题单节点 Hadoop 集群,说白了就是伪分布式环境。我第一次接触这个概念的时候也觉得别扭,明明只有一台机器,为什么叫集群?后来理解了,Hadoop 的伪分布式模式是…

作者头像 李华
网站建设 2026/10/2 2:49:21

零信任微隔离:破解内网横向移动与容器安全的访问控制实战

1. 微隔离到底在解决什么问题先说个我前几年遇到的真实案例。某金融客户内部做了一次攻防演练,红队从一台办公区的跳板机打进了一个测试环境,本来按传统思路,边界防火墙挡得住大部分外部攻击,这就算防线够硬了。可红队进入内网之后…

作者头像 李华
网站建设 2026/10/2 2:47:35

信创文件传输系统选型指南:三条路线与六大指标

1. 先说清楚:为什么政企今年都在聊信创文件传输最近大半年,我身边做政企项目的朋友几乎都被同一个需求找上门:信创文件传输系统。无论是省市级政务云、国企集团、金融机构还是能源单位,招标文件里几乎都有一栏“国产化适配要求”&…

作者头像 李华
网站建设 2026/10/2 2:47:01

SpringBoot3日志实战:Logback配置与MybatisPlus SQL日志排查

系列第02篇,接着上篇搭好的SpringBoot3 MybatisPlus骨架,今天把日志这层彻底补上。日志这东西平时不起眼,真出问题的时候比什么都管用——SQL慢不慢、接口报错在哪、参数到底传成了什么,全得靠它说话。这篇从框架选型讲到logback…

作者头像 李华