ZeRO十年演进:大模型分布式训练显存优化与DeepSpeed实战解析

ZeRO十年演进:大模型分布式训练显存优化与DeepSpeed实战解析 做分布式训练的人2019年之后几乎没人绕得开一个词——ZeRO。十年时间从2015年大家还在为单卡放不下ResNet发愁到2025年动辄千亿万亿参数的基础模型训练稳定跑在多机多卡集群上中间最关键的那个存储器优化方案就是微软DeepSpeed团队提出的ZeRO系列。这个技术解决的是大模型训练里最朴素也最致命的问题显存不够用。它不改变模型结构不改数据集不牺牲精度只是把一张卡装不下的权重、梯度、优化器状态拆开存到多张卡上就能让你用更少的卡训练更大的模型。这篇文章我把ZeRO这十年怎么一步步走到今天的事情捋清楚同时把里面三个Stage的原理、offload到CPU甚至NVMe的玩法、和PyTorch FSDP怎么选、以及我自己用DeepSpeed踩过的一堆坑都写出来。不管你是刚接触大模型训练的新手还是想把自己的训练脚本再压一压显存的老手这篇应该能让你少走不少弯路。1. 从一个显存焦虑的时代说起1.1 2015年的起点分布式训练还停留在“数据并行”2015年做深度学习的人手里拿的模型基本就是AlexNet、VGG、ResNet这个量级几千万到上亿参数。那时候大家说的“分布式训练”默认就是数据并行每张卡复制一份完整模型喂不同的batch然后梯度求平均同步更新。这种方案在当时没有任何问题因为模型本身没多大单卡放得下大家焦虑的是训练速度不是显存。转折来自两个方向。一个是Transformer在2017年横空出世模型规模开始一路狂飙BERT 3亿参数GPT-2 15亿参数GPT-3 1750亿参数。另一个是模型越来越大之后单纯加卡已经失效了——因为每张卡都要放一份完整模型副本显存瓶颈是单卡上限锁死的你加到100张卡单卡放不下还是放不下。我记得2019年调GPT-2大小的模型时一台8卡V100每卡32GB已经要精打细算才能塞进去等到175B的模型出来纯靠数据并行是彻底没戏了。这时候行业里有两条路线开始分叉一条是模型并行/流水线并行把模型的不同层切到不同卡上比如Megatron-LM和GPipe另一条是微软在2019年底放出的ZeRO走的是“状态分区”路线。后来大家都知道了ZeRO及其衍生技术成了训练超大模型的事实标准之一。1.2 一张表看懂显存到底被谁吃掉了先来算一笔账。假设你在训练一个7B参数模型用混合精度FP16参数FP32优化器跑AdamW。一张卡上要装的东西分四块数据对象每个参数占用7B模型总占用说明FP16参数副本2字节14GB前向和反传用的权重FP16梯度2字节14GB反传时累积的梯度FP32主权重4字节28GBAdam更新时用的精确权重Adam动量m4字节28GB一阶矩估计Adam方差v4字节28GB二阶矩估计合计112GB。这是不算激活值、不算临时缓冲、不算通信缓冲的裸数据。一张A100 80GB根本放不下V100 32GB更是想都别想。再加激活值activation7B模型在seq len 2048、batch 4的场景下动辄还要额外几十GB。这就是为什么“单卡放不下”不是一句空话而是明明白白的数学问题。ZeRO做的事情概括成一句话以前是每张卡把上面这堆东西全部复制一份现在是大家合起来平摊各存一份分片用的时候再互相拼起来。1.3 为什么“切一切”就能解决显存问题这里面的直觉可以用仓库来类比。原来每个仓库管理员GPU都自己囤一整箱货完整模型副本地方不够了就堆不下。ZeRO的思路是一个模型的状态不是每次都全量用到的尤其在数据并行场景下每张卡跑的都是同一个模型的同一层只是输入数据不同。既然你们的计算模式完全一样那这些状态完全可以分区保存前向反传到哪一层、需要哪个参数再临时把那一片拉过来。这个“需要时再拼装”的思路就是ZeRO-DP的核心。它的好处是几乎不改变原有的数据并行训练流程不引入复杂的矩阵切分逻辑那是张量并行的事所以工程上落地特别快。2. ZeRO核心设计逐层拆解2.1 三个Stage从只分优化器到连参数都分ZeRO论文里定义了三个递进的Stage对应着分片覆盖的范围逐渐扩大Stage 1P_os只把优化器状态分片。Adam的三块FP32状态主权重、m、v每张卡只存1/N参数和梯度仍然每卡全量持有。这一步已经能省大约75%的优化器相关显存因为Adam那12字节/参数是大头。Stage 2P_osg在Stage 1基础上把梯度也分片。反传过程中梯度是逐步产生的每张卡只保留自己负责的那片梯度然后reduce-scatter汇总。省掉的显存进一步增加。Stage 3P_osgp连模型参数本身也分片。前向传播到某一层时所有卡把这一层参数通过all-gather拼出来算完再丢弃。这是省显存最彻底、通信开销也最大的阶段。拿7B模型在8卡上跑举例用Stage 3每卡只需要存大约112GB 激活值/ 8 ≈ 14GB的静态数据A100 80GB跑起来非常宽裕甚至可以塞更大的batch。这就是很多人对着A100说“8卡能训7B”的底气来源。2.2 通信量为什么没有爆炸一次数学上很划算的交换有人会问参数、梯度、优化器状态都分开了每次都要拼装通信成本是不是高到没法用答案是没有至少没有高到离谱。DDP数据并行每个step要做一次全量梯度的all-reduce通信量大约等于全量梯度的2倍。ZeRO Stage 2用reduce-scatter做梯度归约再用all-gather把更新后的权重广播回去通信量和DDP基本持平但显存占用大幅下降。Stage 3因为参数也要按需gather前向一次、反传一次再加梯度归约通信量大约是DDP的1.5倍左右。换句话说你是用增加的这0.5倍通信量换来了从“单卡放不下”到“随便放”的质变。在大规模集群上这0.5倍通信通常可以通过NVLink、IB、梯度重叠通信来消化绝对收益远大于代价。这也是为什么ZeRO初期最推荐Stage 2它几乎不增加通信压力但已经把优化器状态和梯度这两个最大头解决了8卡训7B、64卡训30B这种场景完全够用。只有模型大到连参数副本都放不下时才需要上Stage 3。2.3 ZeRO-R除了模型状态还有三块隐藏的显存开销ZeRO论文里除了ZeRO-DP之外还专门讲了ZeRO-R负责处理另外三块容易被忽略的开销激活值、临时缓冲、显存碎片。激活值通过“分区激活”partition activations来处理每张卡只存自己负责那部分激活用的时候再跨卡gather。这和重计算activation checkpointing是互补的重计算以2倍前向计算换显存分区激活以通信换显存。两者叠加效果更好也是为什么DeepSpeed配置里经常同时开这两项。临时缓冲主要指all-gather和reduce-scatter的通信缓冲。DeepSpeed的做法是分配一个“恒定缓冲”constant buffer大小可以配置默认可能几百MB到1GB用来避免频繁申请释放造成的不稳定和碎片。显存碎片则通过内存对齐和合并释放来缓解。很多人在Stage 3下遇到莫名OOM其实不是模型太大而是碎片太严重。3. 从DDP到ZeRO的迁移实操3.1 工具选型DeepSpeed还是PyTorch FSDP2023年之后PyTorch原生集成了FSDPFully Sharded Data Parallel效果对标ZeRO Stage 3社区适配也越来越好。所以现在做技术选型时很多人会纠结。我的经验是维度DeepSpeed ZeROPyTorch FSDP功能完整度支持Stage 1/2/3、CPU/NVMe offload、ZeRO主要对标Stage 3offload到CPU较成熟配置复杂度需要写JSON配置灵活但坑多参数化API和PyTorch生态贴合社区资料老牌训练大模型的案例最多更新快PyTorch 2.x之后已成主流与Megatron结合有DeepSpeed-Megatron组合方案需要自己搭桥上手速度中等较快我的建议如果从头写一个新项目PyTorch 2.0以上推荐直接用FSDP代码侵入小如果是要跑成熟的大模型训练框架比如DeepSeek、LLaMA系微调、各类lora训练脚本大概率已经内置了DeepSpeed配置那就直接用DeepSpeed别折腾迁移。两者原理相通很多经验可以互相套用。3.2 一份能直接跑的DeepSpeed配置逐项讲清楚下面是我在单机多卡训7B模型时常用的一份配置按我的经验注释一下{ train_batch_size: 64, gradient_accumulation_steps: 4, optimizer: { type: AdamW, params: { lr: 1e-5, betas: [0.9, 0.95], eps: 1e-8 } }, zero_optimization: { stage: 2, offload_optimizer: { device: cpu, pin_memory: true }, contiguous_gradients: true, overlap_comm: true, reduce_bucket_size: 5e8, stage3_prefetch_bucket_size: 5e8, stage3_param_persistence_threshold: 1e6, stage3_max_live_parameters: 1e9, stage3_max_reuse_distance: 1e9 }, activation_checkpointing: { partition_activations: true } }几个容易踩坑的点stage 2还是37B模型在8卡A100上stage 2已经足够通信快显存占用也稳。只有当你想把batch开得特别大或者卡特别少比如2卡跑7B时才需要stage 3。offload_optimizer到CPU这个配置在stage 2下也能开。它把优化器状态放到系统内存显存省得很猛但训练速度会明显下降因为每次更新都要走PCIe。如果是纯追求训练速度的集群建议别开如果是自己单机调试、卡显存吃紧真香。reduce_bucket_size这个值过小会导致通信频繁切小块过大又占显存。5e8500MB是我常用的平衡点。你可以盯一眼nvidia-smi和NCCL日志再调。activation_checkpointing.partition_activations对应ZeRO-R的激活分区对长序列训练效果非常明显能和重计算叠加。3.3 训练脚本改动清单如果你的脚本之前用的是PyTorch DDP迁移到DeepSpeed只需三步第一命令行启动方式从python -m torch.distributed.launch换成deepspeed --num_gpus8 train.py --deepspeed ds_config.json。第二训练逻辑中把model DDP(model)换成model_engine, optimizer, _, _ deepspeed.initialize(modelmodel, model_parametersparams, configds_config)后续step时用model_engine.step()替代optimizer.step()。第三loss的scaling由DeepSpeed接管不需要你手工做混合精度相关操作除非你自己写了AMP。这里有个很关键的经验DeepSpeed初始化之后原模型的forward/backward接口不变但backward需要传入loss给model_engine.backward(loss)很多人第一次迁移会漏掉这一步导致梯度根本不算。还有如果你原来用torch.cuda.amp.GradScaler迁移后建议删掉否则两个scaler会互相打架出现莫名其妙的loss不下降。3.4 容量估算上集群前先算一笔账上大规模集群前我强烈建议先做一次显存估算。以7B为例6张A100 80GB、跑stage 2、不开offload每卡可用显存里静态占用大约30GB左右剩下的空间可以用来放激活值和更大的batch。反过来如果你的模型是30B8卡stage 3起步激活值一定要开重计算offload可以视带宽情况决定。一个粗略公式经验值非精确Stage 2每卡占用约为FP16参数 FP16梯度 Adam状态/N再乘1.2的安全系数。Stage 3每卡占用约为FP16参数 FP16梯度 Adam状态/ N再乘1.3。以7B、8卡为例stage 2约141484/8×1.2 ≈ 42GBstage 3约112/8×1.3 ≈ 18GB。算出来的数如果已经接近显存上限先砍激活值、砍batch再考虑开offload不要一上来就把所有手段都堆上。4. 十年路线图从Stage 1到ZeRO三代4.1 2015-2025关键时间线要理解ZeRO这十年不能只看它2019年横空出世那一瞬间。我把这条线按年份拉出来2015-2017数据并行模型并行并行着走显存问题开始显现但还没到不可收拾。2018-2019GPT-2/XLNet大规模模型出现Megatron-LM和张量并行成为热点。微软研究院启动ZeRO项目目标是“任何人用任何GPU都能训任意大的模型”。2019年底ZeRO论文上线提出ZeRO-DP和ZeRO-R首次在论文层面实现“train 100B模型”的理论可行性。2020年初DeepSpeed开源附带ZeRO Stage 1/2实现。社区第一次可以在普通多卡机器上跑BERT-Large、GPT-2 1.5B成本骤降。2021年ZeRO-Offload发布把优化器状态和梯度卸载到CPU内存随后ZeRO-Infinity扩展到NVMe最大亮点是单卡也能训练超大模型。2023年ZeRO第一代发布引入分区通信partitioned communication、量化通信quantized communication和低精度参数解决Stage 3通信瓶颈。2024-2025年ZeRO第二代/第三代迭代低精度优化器状态、混合精度的通信、对推理场景的ZeRO-Inference支持逐步落地。同时PyTorch FSDP全面成熟ZeRO思想被吸收进各家框架。“ZeRO十年演进”严格说ZeRO本身是2019年才诞生的但2015年这个起点非常适合理解它解决的是什么时代的问题显存从“不够大的烦恼”变成了“决定模型规模的硬边界”。4.2 ZeRO-Offload和ZeRO-Infinity把显存边界推到CPU和硬盘ZeRO-Offload的思路非常朴素既然GPU显存不够CPU内存通常大得多服务器轻松512GB到1TB那就把不经常用、但占地方的东西放到CPU上。它的最佳实践是优化器状态放CPU梯度也放CPU参数留在GPU反向传播和参数更新在CPU上异步执行。实测下来对于几十B级别模型CPU offload能让人用少量GPU卡跑起来代价是速度明显下降。后来ZeRO-Infinity把这一套扩展到了NVMe SSD允许优化器状态和梯度直接落盘这就把“单卡能训的模型上限”推到了接近无限大——理论上只要你有足够的CPU内存和硬盘单卡也能训170B参数模型只是速度慢到怀疑人生。这个方案适合调试、临时跑通不是生产环境的首选。实际使用中offload到CPU的配置我之前已经给了NVMe offload要再加一层offload_optimizer: { device: nvme, nvme_path: /mnt/nvme, buffer_count: 4, fast_readwrite: true }注意nvme_path必须是一个真实的本地NVMe设备路径不能是网络盘否则延迟会高到完全跑不动。fast_readwrite开启后DeepSpeed会为每个buffer分配一块固定内存做DMA对这种场景提升很大。4.3 ZeRO三代通信量的正面硬刚ZeRO Stage 3最大的痛点是通信量是DDP的1.5倍所以在千卡集群上通信优化比显存优化更能决定训练速度。ZeRO就是冲着这个去的。第一代ZeRO在2023年放出三个关键机制分区分组通信把大world size拆成小分组分层all-gather、量化通信把通信数据从FP16降到INT8传输量减半精度损失通过量化感知的梯度压缩来弥补、低精度优化器状态直接砍掉一部分优化器状态的存储需求。第二代和第三代继续在通信分组策略、混合精度调度以及GPU/NVMe分层offload上做文章。看起来复杂落地其实很透明DeepSpeed的配置里打开comms_logger和zero_quantized_ gradients相关开关或者直接用启动参数--zero-stage3 --zero-quantized-gradients。实测在100Gbps网络的多机场景下量化通信能带来20%-40%的端到端提速。但是注意如果你的网络带宽已经很充裕比如全部NVLink互联这个优化收益不大反而可能因为量化反量化增加GPU计算。4.4 和Megatron、FSDP到底什么关系一张图理清这里我经常被人问ZeRO和Megatron不是冲突吗不是。它们负责的维度完全不同方案切分维度适合场景ZeRO-DPstage 3数据并行中的模型状态分片同层参数多模型大但单层能装下集群通信条件好Zhang量并行Megatron按矩阵维度切分每层参数单层矩阵过大如超大hidden size流水线并行PP按层切分层数特别多机器间带宽一般FSDP和ZeRO-DP类似PyTorch生态内想快速实现实践中训练一个175B模型典型组合是数据并行ZeRO Stage 2或3张量并行8路流水线并行若干阶段激活重计算。单一方案解决不了所有问题ZeRO负责的是“让每张卡不存冗余状态”这部分和Megatron的张量切分是互补的。5. 实战中的坑与排查速查5.1 六个高频报错/现象及处理建议OOM已经在Stage 3了还是显存爆了。优先检查三件事reduce_bucket_size是不是太大激活值有没有开重计算和分区offload_param/offload_optimizer是否真的生效可以在DeepSpeed启动日志里确认。A100上直接看Deepspeed info输出它会打印每块显存分配明细。训练卡死或异常慢NCCL超时。多机场景最常见通常和网络环境有关容器里没正确指定NCCL_SOCKET_IFNAME或者NCCL_P2P_DISABLE设置不对。我一般先设NCCL_DEBUGINFO看卡在哪一次集合通信上再逐项调。加载checkpoint后loss震荡或完全乱掉。一般是optimizer state没有正确加载或者resume_from_checkpoint路径给错DeepSpeed只会从checkpoint目录里的zero_pp_rank_*.pt恢复确保路径统一。推理或agent调用时报unknown model。一些模型服务或agent框架在启动时会把模型名映射到实际权重文件如果填的模型名和注册名不一致就会出现类似unknown model: xxx的报错。这跟ZeRO本身关系不大多半是配置里模型的name、路径、版本号没对齐。查的时候先看配置里model_name_or_path和框架自带的模型注册表别一上来就怀疑显存。开启offload后训练变慢到无法接受。这是“显存换速度”的必然代价。建议只offload优化器不要连参数也offload开pin_memory确保数据加载不吃CPU有NVMe的话优先用NVMe而不是SATA SSD。梯度更新前后loss完全不动。检查你是不是把model_engine.backward(loss)写成了loss.backward()以及有没有在调用step之前手动调用了zero_grad。DeepSpeed的step内部会处理梯度清空手动zero_grad有时会把累积的梯度清掉。5.2 调优的几个“土办法”这些不是我编的是每次新上一个训练任务我都会在日志里用一组固定手段观测先看nvidia-smi里显存和功耗。显存用了90%以上但功耗只有一半说明通信或数据加载瓶颈不是计算瓶颈。看NCCL日志里的带宽数字。如果远低于预期比如100Gbps网络上只有20Gbps先查是不是走错了网卡。开overlap_comm后再看吞吐如果提升不明显说明通信本来就不是瓶颈别再做量化通信了。用小batch跑通再逐渐放大找到显存和吞吐的平衡点。这一步能帮你判断是静态显存瓶颈还是激活值瓶颈。5.3 搜索时别把ZeRO和那些“zero”搞混了写这篇文章时我顺手搜了一下发现现在搜ZeRO出来的内容会被各种其他项目稀释嵌入式里的荔枝派Zero、Go生态里的go-zero框架、无人机里的Zero Omega、甚至一些AI agent工具报错里的unknown model提示。它们和深度学习显存优化完全是两码事。想查技术资料时建议关键词带上下文集ZeRO DeepSpeed、ZeRO stage 3、ZeRO offload或者直接去DeepSpeed官方文档和论文源码里找能省不少时间。6. 我的一些个人体会回头去看这十年ZeRO最值得学习的不是某个具体的显存优化技巧而是“发现问题、量化问题、系统解决问题”的思路。它没有发明新的并行范式而是在已有数据并行框架里把“冗余存储”这四个字抠到了极致然后把通信成本控制在一个可接受的范围内最终让大规模训练从“大厂专用”变成了“工程师可及”的能力。我自己的实践体会是拿到一个训练任务不要一上来就无脑Stage 3 offload先从Stage 2骑一遍用nvidia-smi看真实显存分布再结合模型大小和卡数决定要不要升级Stage 3。多数时候Stage 2就够用而且调参成本低得多。真正上到Stage 3时一定先把通信环境测好、NCCL设置调好否则你会被各种超时和卡死折磨到怀疑人生。最后再分享一个小技巧在DeepSpeed训练脚本里加一行torch.cuda.memory_summary()跑两个step后输出显存分配详情你能看到哪些buffer占了大头。很多看起来玄乎的OOM一查就原形毕露。这比在网上盲搜报错要靠谱得多。