1. 为什么大模型训练总被显存卡脖子
老读者应该知道,我有个习惯:遇到好技术,不满足于“会用”,总想把论文翻来覆去读几遍,拆开揉碎,搞明白它到底解决了什么问题、为什么这样设计。今天要细读的,是DeepSpeed体系里最经典的一篇——ZeRO(Zero Redundancy Optimizer)。
先说个背景。这几年大模型规模膨胀得太快,早些时候训练个十亿参数模型,单卡A100 80G还能勉强塞进去。到了百亿、千亿参数,单卡别说训练了,光把模型参数放进去都费劲。很多人第一反应是上模型并行或者流水线并行,但这两兄弟实现复杂、通信模式也重,对绝大多数团队来说门槛太高。ZeRO的出现改变了这个局面——它让你在数据并行的框架下,就能训起远超单卡显存上限的模型,而且改动量小得惊人。
这篇文章写给谁?正在做LLM微调或预训练、被OOM折腾得想砸机器的人;想用多卡但不知道显存是怎么满的、该怎么省的人;以及想真正理解DeepSpeed配置里那些stage参数背后含义的工程师。我会从显存账单讲起,把ZeRO的Stage 1到Stage 3、Offload和ZeRO++挨个拆开,最后附上实战配置和踩坑记录,保证你能直接抄作业。
先做一道简单的算术题:训练一个1.5B参数的GPT类模型,用混合精度(fp16)在单卡上跑,显存到底消耗在哪?一共三块:参数,fp16下约3GB;梯度,跟参数同尺寸,也是3GB;优化器状态,这里才是大头——Adam要维护一份fp32的master权重副本、一份一阶动量、一份二阶动量,每个都是4GB,加起来12GB。这三类合计18GB,而其中优化器状态占了三分之二。跑过训练的人看到这个数字应该不陌生:模型参数才3GB,怎么显存一下就用掉20多GB?就是这些“隐形状态”在作祟。
这还只是1.5B模型。换到175B的GPT-3,fp16参数就要350GB,梯度350GB,Adam的三份fp32状态各700GB、合计2100GB。一个峰值收敛到两千多GB的显存需求,单卡根本无解。所以你缺的不是模型显存,是“训练过程显存”。
2. 先看Data Parallelism的冗余问题
2.1 DDP是怎么工作的
在ZeRO出现之前,大家最常用的多卡训练方案是DDP(Distributed Data Parallel)。它的逻辑很朴素:每张卡上都放一份完整的模型参数、梯度和优化器状态,各自喂不同的batch数据,前向反向算完,再用一次All-Reduce把梯度跨卡求和并回传,之后每张卡各自跑优化器更新。
这种方案的好处是简单、通信效率高、对模型结构零侵入。坏处也一眼就能看见:N张卡,就意味着同一份参数、梯度、优化器状态被复制了N份。8卡训练1.5B模型,显存需求直接从18GB膨胀到144GB,绝大部分都是重复的。DDP的前提假设是“单卡能装下整个训练状态”,一旦模型变大,这个前提就崩塌了。
有人会问:模型并行不是能解决吗?能,但模型并行(如Megatron-LM的张量并行)要把每层的矩阵切到不同卡上,需要卡间紧密同步,通信量巨大,而且对代码改动很重。流水线并行虽然按层切分,但存在空泡问题、微批次调参繁琐。对大多数场景来说,这些都是“牛刀”,杀鸡不合适,杀牛也费劲。
2.2 ZeRO的切分哲学
ZeRO的选择很有意思——继续保持数据并行的框架,不切模型、不搞流水线,而是把DDP里冗余的三类状态切开来,让每张卡只保存一部分。参数、梯度、优化器状态都不再是每卡全量,而是按rank均匀切分:原本每张卡保存一份完整副本,现在N张卡合起来保存一套完整状态,单卡承担1/N。
这中间有个关键点:ZeRO是用通信换显存。虽然不存全量了,但前向计算时要用到完整的参数、反向要用到完整的梯度,那就需要在合适的时机通过通信手段把东西“聚齐”。读论文的时候你会发现,ZeRO的设计美感就在于它把“什么时候需要什么东西”拿捏得很准,三者状态的分片策略逐级递进,对应Stage 1、Stage 2、Stage 3。
我用个类比帮助理解:以前DDP就像团队里每个人各背一个装满全套工具的大工具箱,走路慢、重复又重;ZeRO则是把这套工具拆开,你带螺丝刀、我带扳手、他带钳子,开工时谁要工具就现场喊一声借过来,用完还回去。付出的代价是喊话和传递的时间,省下的却是每个人背包的重量。在大模型场景下,“背包重量”是实实在在的显存下限,这个交换非常划算。
3. ZeRO核心:三种状态逐级优化
3.1 Stage 1:切分优化器状态
Stage 1的思路是:参数和梯度的存储仍然每卡全量保留,只把优化器状态按参数维度切分到各卡。还是用linearly切分方式,比如64卡就跑64路数据并行,每个rank只维护1/64的Adam状态。这样1.5B模型的优化器状态从12GB变成12GB/64 = 0.1875GB,再加上参数和梯度各3GB,单卡占用量约6.2GB。对更大模型提升更明显,因为优化器状态本来占大头。
但这里有个执行层面的问题:传统DDP里每张卡算完全量梯度后各自调优化器更新全部参数,现在优化器状态切分了,每张卡只负责更新自己那一份参数对应的梯度。所以反向传播阶段不能再用All-Reduce来同步全量梯度,改用Reduce-Scatter——先做跨卡归约,再把结果按rank切分,每卡只拿到自己负责的那一块片段。这样每个rank只计算并保存与自己的分片对应的梯度值,然后本地更新自己那部分参数。
注意一点:参数本身仍然是全量保存在每张卡的,更新完自己分片的参数后,需要通过一次All-Gather把新参数广播给所有卡。所以Stage 1的训练循环里,每轮有一个Reduce-Scatter加一个All-Gather,通信量和DDP的All-Reduce差不多,但显存占用大幅下降。
3.2 Stage 2:切分梯度
Stage 2想再往前走一步:既然Stage 1里每卡只负责更新一部分参数,那更新之前本来只需要接收对应的那一小块梯度,为什么还要在全量梯度上算一遍?于是ZeRO把梯度也切分了。注意,这里切分的是“在通信和存储过程中的梯度”,不是反向传播算法本身。反向传播时,每层梯度算出来后,streaming地把不属于当前rank的梯度分片释放掉或发给对应rank,最终每卡只保留自己负责的那1/N梯度。
这一刀下去,1.5B模型的梯度从3GB降到0.047GB,加上参数3GB和分片后的优化器状态0.1875GB,单卡占用约3.2GB左右。跟DDP的18GB相比,已经省了超过80%。Stage 2在通信模式上跟Stage 1基本一致,仍然是一轮Reduce-Scatter加一轮All-Gather,通信量没有额外增加,但显存又省了一大块。这也是为什么业界普遍把Stage 2当成性价比最高的默认选项。
3.3 Stage 3:连参数也切了
Stage 3是完整形态:参数、梯度、优化器状态全部切分。到这里,代价开始显现,因为前向和反向过程中每个算子都需要用到完整的参数,而参数不再全量保存在本地。ZeRO的做法是,在前向过程中按需All-Gather去临时取回当前层或当前分段的完整参数,计算完成后再丢弃本地非分片部分。反向同理,需要重新取回参数来计算该层的梯度。
带来的显存收益也非常惊人。1.5B模型在Stage 3下单卡只需要约0.2GB左右的参数分片、0.047GB梯度、0.1875GB优化器状态,算上激活值等杂项,通常能把整个训练状态压在1GB以内。在很多对比实测里,用ZeRO Stage 3配合高速NVLink,甚至可以在8张V100上训练17B参数的模型——这在DDP下是不可想象的(DDP需要每张卡单独存下17B的全部训练状态)。代价是通信量显著增加:每轮前向和反向各需要一次全量参数的All-Gather,轮末还要Reduce-Scatter梯度,总通信量几乎是DDP的1.5倍左右。通信开销上去了,但显存天花板被彻底打开了。
下表总结了三个Stage的显存需求和通信变化:
| 方案 | 优化器状态 | 梯度 | 参数 | 通信量 | 适合场景 |
|---|---|---|---|---|---|
| DDP | 每卡全量 | 每卡全量 | 每卡全量 | 每轮约2倍参数量 | 单卡能放下全量训练状态 |
| ZeRO Stage 1 | 切分 | 每卡全量 | 每卡全量 | 约等于DDP | 优化器状态是唯一瓶颈 |
| ZeRO Stage 2 | 切分 | 切分 | 每卡全量 | 约等于DDP | 显存紧张但通信带宽中等 |
| ZeRO Stage 3 | 切分 | 切分 | 切分 | 约1.5倍DDP | 追求极限显存容量,需要高速互联 |
4. 显存还不够?Offload和ZeRO++
4.1 把状态搬到CPU:ZeRO-Offload
Stage 3已经切得很极致了,但如果模型实在太大、GPU卡数量又有限,仍可能放不下。这时还有个方向:把训练状态的一部分搬到CPU内存甚至NVMe磁盘上去。这就是ZeRO-Offload。
另辟蹊径的思路是这样的:GPU显存贵且有限,CPU内存便宜又大,一张A100 80G的显存价格能买好几台大内存机器。ZeRO-Offload最经典的组合是Stage 2加上优化器状态Offload:优化器状态全部放到CPU内存,GPU只保留参数、梯度和前向/反向计算所需的临时内存。每轮迭代,GPU算完梯度后把更新所需的数据传到CPU,CPU跑Adam更新,再把更新后的参数传回GPU。这样的设计把12GB里的大头(优化器状态)直接挪走,1.5B模型GPU显存需求能压到接近3GB的水平。
但要注意,Offload的核心瓶颈是PCIe带宽。CPU内存和GPU之间靠PCIe传输,带宽通常只有几十GB/s,跟显存内部几个TB/s比差了两个数量级。如果训练过程频繁搬动全部参数,速度会被拖垮。实操里要把“哪个状态放哪里”平衡好,也需要看懂DeepSpeed的配置参数写的CPU offload优化器状态而不是全部offload。我的经验是:Offload适合追求“能跑起来”的探索场景,比如单机双卡训练10B级别模型;真要大规模生产训练,还是优先加GPU卡数,让数据并行数上去之后靠Stage切分解决问题,不太依赖Offload。
4.2 针对通信和滞后的优化:ZeRO++
ZeRO++ 名字听起来像加强版,解决的是我在前面提到的痛点:Stage 3通信量偏大,Offload又受限于PCIe带宽。它给出了几个关键改进:
第一是对通信权重做量化。前向All-Gather参数时,本来传fp16的2字节权重,ZeRO++把它压到INT8甚至更低精度,只在本地乘算之前做反量化还原。这相当于把通信数据体积直接砍半,带宽压力立刻小很多。代价是有损压缩,需要精调量化策略,但在大规模训练里这个取舍通常值得。
第二是分层参数分区。ZeRO++意识到一个模型里不同层被调用的频率完全不同,不需要每次All-Gather都全量广播。它把参数按层分成主分片和从分片:高频访问的层在本机或本地节点保存完整副本,低频访问的层做全局切分。这套设计能显著减少跨机通信。
第三是offload路径上的NVMe优化。ZeRO++里的Infinity机制可以把参数、优化器状态进一步搬到NVMe SSD上,配合异步预取、流水线传输,用大容量闪存换取极致显存弹性。我的看法是,ZeRO++更适合“算力充裕但互联带宽有限”的集群——比如买了很多GPU但交换机上不了200Gbps的小团队,通过量化通信能实打实压下来传输时间。
5. 实战配置:从接入DeepSpeed到踩坑记录
5.1 一份能直接用的DeepSpeed配置
写代码之前,我先举个例子。假设你用HuggingFace Transformers训练一个7B模型,DeepSpeed的接入方式通常是在训练脚本里加deepspeed.init_distributed()和TrainingArguments(deepspeed="ds_config.json")。核心的ds_config.json长这样:
{ "train_batch_size": 32, "gradient_accumulation_steps": 2, "optimizer": { "type": "AdamW", "params": { "lr": 3e-5, "betas": [0.9, 0.999], "eps": 1e-8, "weight_decay": 0.01 } }, "zero_optimization": { "stage": 2, "allgather_partitions": true, "reduce_scatter": true, "contiguous_gradients": true, "offload_optimizer": { "device": "cpu", "pin_memory": true } }, "fp16": { "enabled": true, "loss_scale": 0, "loss_scale_window": 1000, "initial_scale_power": 16 }, "gradient_clipping": 1.0 }几个关键点:stage直接决定用哪个阶段的ZeRO;allgather_partitions和reduce_scatter控制通信算子是否启用分区聚合;contiguous_gradients把零散的梯度张量合并成连续缓冲区,减少内存碎片;offload_optimizer只有在显存实在不够时才开。我建议的参数启动顺序是这样:先开Stage 2,能跑就不要动;再尝试把offload_optimizer关掉,避免PCIe瓶颈;最后实在放不下再上Stage 3。
5.2 如何选择Stage:一张实操决策表
很多人问到底选哪个Stage,我给个简单决策树:
- 单卡显存能放下全部训练状态:直接DDP,ZeRO反而增加通信逻辑;如果只是优化器状态超了,上Stage 1。
- 多卡训练但每张卡的显存只有峰值需求的一半到三分之一:Stage 2是首选,通信开销跟DDP几乎持平。
- 显存相差一个数量级,比如训练几十B模型只有8卡40G:必须Stage 3,同时确认NVLink或高速网络,否则训练会卡在通信上。
- 单机单卡但想跑超大模型:Stage 2 + Offload最稳,别选Stage 3,因为没有多卡做分片,Stage 3不会带来额外收益反而增加通信。
这里补充一个重要认知:ZeRO解决的是峰值显存占用,不会改变训练动态本身。也就是说,在不考虑通信开销的理想情况下,ZeRO Stage 3训练出的模型和DDP在数学上是等价的(同样的优化器更新序列),这也是它能成为通用方案的原因——不会因为引入分区就影响收敛质量。
5.3 我踩过的坑和排查心得
用ZeRO这两年,我在实际训练中踩过不少坑,挑几个典型的说说。
第一个坑是梯度累积和batch size的配比。ZeRO开启后,train_batch_size = per_gpu_batch_size × world_size × gradient_accumulation_steps。很多人把gradient_accumulation_steps设得太大,结果loss剧烈波动——因为梯度是在分片状态下累积的,跨step的梯度累积跟DDP语义虽然一致,但浮点累积顺序变化导致数值微小漂移。这个不用纠结,训练过程正常即可。但注意learning rate warmup要重新调,因为有效batch变大后收敛节奏会变。
第二个坑是动态loss scale和梯度裁剪。混合精度训练会维护一个动态loss scale,每次迭代如果发生overflow就缩小,一段时间内没overflow就放大。ZeRO的分区梯度模式下,overflow检测需要跨卡同步,DeepSpeed会收集所有rank的scale状态再做全局归一。如果代码里自己手动改了梯度裁剪逻辑,有几张卡梯度被clip、有几张没有,就会出现优化器更新不一致,训练直接发散。我一直用DeepSpeed内置的gradient_clipping参数,不要自己在training loop里手动clip。
第三个坑是contiguous memory buffer导致的显存虚增感。开了contiguous_gradients后,显存使用曲线会一下子跳高一大块,这是分配了固定连续缓冲区,不是泄漏。很多人在NVIDIA SMI里看到显存90%就慌了,实际训练初期看起来高占用很正常。关掉这个参数可以省点显存,但会损失训练吞吐,不推荐关。
第四个坑是关于checkpoint的保存。ZeRO Stage 3的模型权重是分片存储的,直接torch.save(model.state_dict())保存出来的是分片权重,合并回单卡模型需要自己写逻辑。我的建议是启用DeepSpeed的save_16bit_model或用HuggingFace的zero_to_fp32.py脚本,把分片合并成完整的bf16/fp16权重,再做后续转换。教训是:一开始我图省事直接保存分片,结果下游推理时加载模型怎么都不对,白折腾了半天。
再分享一个小技巧:开启overlap_comm和reduce_scatter组合。DeepSpeed在Stage 2上有一系列通信重叠优化。反向传播的梯度reduce可以和下一层的前向计算并行起来,虽然配置项只改overlap_comm: true,但实测很多场景吞吐能提升10%-20%。如果训练速度上不去,优先检查这个而不是怀疑代码写错了。
6. 影响与未来:ZeRO带来了什么
ZeRO的影响已经远超DeepSpeed本身。OpenAI和Microsoft的早期论文里,175B模型的训练之所以在千卡级别GPU集群上可行,背后就有ZeRO撑着。后来PyTorch官方推出的FSDP(Fully Sharded Data Parallel)本质上就是复刻了ZeRO Stage 3的思维——全参数分片、按需All-Gather、计算完丢弃。换句话说,ZeRO的设计思路已经成了大模型训练的基础设施级标配。
对普通工程师来说,ZeRO最大的价值是降低了并行训练门槛。以前训不动的大模型,现在改几行配置文件就能跑;以前只能单卡跑的小模型,现在几张卡就能扩展。它没有改变优化算法本身,却通过巧妙的存储布局改变了显存的物理边界,这种“不动算法只动系统”的思路很值得借鉴。
从我的实践经验看,ZeRO对集群通信的要求其实也没那么可怕。单机8卡内用NVLink灌满基本无压力;跨机训练时,只要节点间带宽不低于40Gbps,Stage 2日常很稳。真要上Stage 3做超大模型,才需要认真规划网络拓扑和通信分桶策略。这份细读写到最后,核心是让更多人理解:显存不是决定模型规模上限的唯一门槛,你的调度和工程优化能力同样是关键变量。