DeepSpeed_ZeRO
2026/7/23大约 5 分钟
DeepSpeed ZeRO 是由微软(Microsoft)开发的一项在大模型分布式训练中具有里程碑意义的显存优化技术。它的全称是 Zero Redundancy Optimizer(零冗余优化器)。
在它诞生之前,传统的 DDP(数据并行) 训练非常简单粗暴:每张 GPU 都必须在显存里完整复制一份一模一样的模型参数、梯度和优化器状态。当模型大到一定程度时,这种“全量复制”会导致显存瞬间炸裂(OOM)。
ZeRO 的核心哲学极其务实:既然大家都在同一个集群里,为什么每张卡都要存一份完整的数据?我们为什么不把这些显存里的狗皮膏肉像切蛋糕一样平均分片,等需要算哪一层时,再通过 NCCL 临时向别人要?
为了把显存榨干,ZeRO 提出了著名的三阶段剥皮法(Stage 1/2/3),每升一级,显存省得越多,但网络通信的压力就成倍暴涨:
一、 ZeRO 的三阶段显存深度解剖
大模型在训练时,显存里主要装了三大主体(业界统称为 Model States):
- Optimizer States (优化器状态):比如 AdamW 优化器,占了近 60% 的绝对大头。
- Gradients (梯度):反向传播算出来的修正量。
- Parameters (模型权重/参数):模型本身的本体。
Stage 1:只切分“优化器状态” (Optimizer States Sharding)
- 手法: 显存里最胖的优化器状态不再在每张卡上完整保留,而是平均切成 $N$ 份($N$ 为总卡数)。每张卡依然保留完整的模型参数和梯度。
- SRE 网络视角: 更新权重时,由于每张卡只负责更新自己对应那 $1/N$ 的参数优化,算完后需要调用一次
NCCL AllGather,把更新后的全量参数同步给所有人。 - 回报率(ROI): 极高!能省下近 4 倍的优化器显存,而网络通信量几乎没有变多,非常稳健。
Stage 2:连“梯度”一起切分 (Gradients Sharding)
- 手法: 在 Stage 1 的基础上,反向传播算出来的梯度也不在每张卡上全量保留了。谁负责更新哪部分参数,当某层算完梯度后,不属于它的梯度直接扔掉。
- SRE 网络视角: 反向传播每算完一层,不再走 DDP 经典的
AllReduce(全量求和并同步),而是走NCCL ReduceScatter(求和但只把碎片分给对应的卡)。网络通信量和 DDP 完全一模一样,但显存又省了一大块。
Stage 3:连“模型参数”也切分 (Parameters Sharding) —— 终极榨汁机
- 手法: 丧心病狂。每张卡连完整的模型参数都不存了!整张卡里空空如洗,只存了 $1/N$ 的模型参数碎片。
- SRE 网络视角(极度吃跨机网卡带宽):
- 前向传播时:算到第一层,所有卡通过
NCCL AllGather瞬间把第一层的参数拼完整,算完,立刻把别人的参数从显存里抹除。再算第二层,再拼,再抹除…… - 反向传播时:同样的操作再来一轮。
- 评价: 只要你卡足够多,你能跑得下无限大的模型。但你的 MFU(算力利用率) 会被高频的
AllGather网络通信直接拉垮,除非你配了无损的高速 RDMA 网络。
二、 还有一招保命技:ZeRO-Offload (异构内存切分)
如果你在配置 DeepSpeed 的 JSON 文件时,看到了下面的参数,这就是大名鼎鼎的 Offload 技术:
"offload_optimizer": {
"device": "cpu"
}
- 本质: 借用“后勤力量”。如果你们公司的 GPU 显存实在是小(比如要在单机 8 卡上强行微调一个很大的模型),显存连 Stage 3 都撑不下,ZeRO 允许你把最胖的优化器状态(甚至梯度)直接从显存里吐出来,扔进系统的 CPU 内存(甚至外部硬盘)里。
- SRE 避坑预警: 这一招能绝对保证你的代码不报
CUDA Out of Memory。但是!因为数据要通过 PCIe 总线在 CPU 和 GPU 之间疯狂来回搬运,会让你训练的吞吐量(Tokens/s)暴跌数倍甚至 10 倍。在工业级生产中,除非万不得已,否则尽量不要开。
三、 ZeRO 与 PyTorch FSDP2 的爱恨情仇
你可能会问,我们前面一直在学 Meta 的 FSDP2,它和微软的 ZeRO-Stage 3 听上去不是一模一样吗?
它们在哲学目标上是一致的(参数、梯度、优化器全分片),但在底层实现上,FSDP2 代表了更先进的生产力:
- 底层抽象不同: * ZeRO-Stage 3:是一个外挂式的第三方框架。它为了方便切分,会把模型所有的矩阵强行拉直、打平成一锅 1D 的大数组(FlatParameter)。这导致它在训练中途保存 Checkpoint(模型快照) 时极度痛苦,经常会把 CPU 内存撑爆(OOM),且很难和张量并行(TP)完美融合。
- FSDP2:是 PyTorch 原生发起的。它基于 DTensor 架构,完全保留每一层矩阵原汁原味的形状,并能利用自动的异步预取机制(Overlap),把
AllGather通信完美藏进计算时间的阴影里。
SRE 实战选型指南
- 如果你们算法组的微调代码是基于 Hugging Face 社区、或传统的 Transformers 库 + DeepSpeed 脚本,那就顺水推舟用 ZeRO-Stage 1 或 2(性价比极高)。
- 如果你们正在基于最新的 TorchTitan 或者是纯原生 PyTorch 2.x 从零预训练一个千亿大模型,毫无疑问,直接上 FSDP2 / HSDP,它的算力回报率(MFU)和保存快照的稳定性会明显优于 DeepSpeed Stage 3。
