当前位置: 首页 > news >正文

【硬核拆解】DeepSpeed ZeRO:从56GB到7GB,三阶段分片如何让大模型训练显存暴降87.5%?

目录

  1. DeepSpeed ZeRO 的设计动机
  2. ZeRO-1:优化器状态分片
  3. ZeRO-2:梯度分片
  4. ZeRO-3:全参数分片
  5. ZeRO-Offload 与卸载
  6. DeepSpeed ZeRO 的边界与失效模式

摘要

DeepSpeed ZeRO(Zero Redundancy Optimizer)通过分阶段消除数据并行中的冗余存储,将显存占用降低到原来的 1/N。ZeRO-1 分片优化器状态,ZeRO-2 分片梯度,ZeRO-3 分片全部参数。本文从 ZeRO 的设计动机出发,分析三阶段的分片原理、通信模式和卸载策略。

1. DeepSpeed ZeRO 的设计动机

数据并行训练中,每个 GPU 持有完整的模型参数、梯度和优化器状态副本。这些副本是冗余的——每个 GPU 上的参数值完全相同。ZeRO 的核心思想是:消除冗余存储,只在需要时收集完整数据

1.1 数据并行的冗余分析

存储内容每个 GPU 存储实际需要冗余度
模型参数完整(14GB for 7B)分片(14GB/N)N
梯度完整(14GB for 7B)分片(14GB/N)N
优化器状态完整(28GB for 7B, Adam)分片(28GB/N)N
总计56GB56GB/NN

1.2 ZeRO 的核心思想

ZeRO 的核心思想是分阶段消除冗余

DDP 冗余存储

ZeRO-1: 分片优化器状态

ZeRO-2: 分片梯度

ZeRO-3: 分片参数

显存节省: 4x (Adam)

显存节省: 8x

显存节省: 12x

1.3 DeepSpeed ZeRO 的历史演进

ZeRO 论文(2019)→ ZeRO-1/2 实现(2020)→ ZeRO-3 全分片(2020)→ ZeRO-Offload(2021)→ ZeRO-Infinity(2022)。

1.4 DeepSpeed ZeRO 的产业应用

模型规模ZeRO 阶段GPU 数
BERT-Large340MZeRO-264
GPT-3175BZeRO-310,000
LLaMA 65B65BZeRO-32,048
BLOOM 176B176BZeRO-3384

1.5 DeepSpeed ZeRO 的局限性

ZeRO 的局限性包括:通信量增加(分片越多,通信量越大)、实现复杂度高(需要手动管理分片)以及小模型收益有限(小模型下 ZeRO 的收益不如 DDP)。

2. ZeRO-1:优化器状态分片

2.1 ZeRO-1 的原理

ZeRO-1 只分片优化器状态,模型参数和梯度保持完整。优化器状态(如 Adam 的动量和方差)占显存最大(通常是模型参数量的 2 倍),分片后显存节省显著。

2.2 ZeRO-1 的显存节省

分片内容未分片(7B, FP16)分片后(8 GPU)节省
模型参数14GB14GB0%
梯度14GB14GB0%
优化器状态28GB3.5GB87.5%
总计56GB31.5GB43.75%

2.3 ZeRO-1 的通信

ZeRO-1 在优化器更新时需要通信:每个 GPU 只更新自己的分片,然后通过 All-Gather 收集完整更新后的参数。

2.4 ZeRO-1 的实现

importdeepspeed# ZeRO-1 配置zero_config={"zero_optimization":{"stage":1,# ZeRO-1"reduce_bucket_size":5e8,"allgather_bucket_size":5e8}}model_engine,optimizer,_,_=deepspeed.initialize(model=model,optimizer=optimizer,config_params=zero_config)

3. ZeRO-2:梯度分片

3.1 ZeRO-2 的原理

ZeRO-2 在 ZeRO-1 的基础上,进一步分片梯度。每个 GPU 只存储本分片参数的梯度,不存储完整梯度。

3.2 ZeRO-2 的显存节省

分片内容未分片(7B, FP16)分片后(8 GPU)节省
模型参数14GB14GB0%
梯度14GB1.75GB87.5%
优化器状态28GB3.5GB87.5%
总计56GB19.25GB65.6%

3.3 ZeRO-2 的通信

ZeRO-2 在反向传播时使用 Reduce-Scatter 分发梯度,在优化器更新后使用 All-Gather 收集参数。

3.4 ZeRO-2 的实现

# ZeRO-2 配置zero_config={"zero_optimization":{"stage":2,# ZeRO-2"reduce_bucket_size":5e8,"allgather_bucket_size":5e8,"contiguous_gradients":True,"overlap_comm":True# 通信重叠}}

4. ZeRO-3:全参数分片

4.1 ZeRO-3 的原理

ZeRO-3 在 ZeRO-2 的基础上,进一步分片模型参数。每个 GPU 只存储本分片参数,不存储完整参数。

4.2 ZeRO-3 的显存节省

分片内容未分片(7B, FP16)分片后(8 GPU)节省
模型参数14GB1.75GB87.5%
梯度14GB1.75GB87.5%
优化器状态28GB3.5GB87.5%
总计56GB7GB87.5%

4.3 ZeRO-3 的通信

ZeRO-3 在前向和反向传播时都需要 All-Gather 收集完整参数,计算后丢弃非本分片参数。

4.4 ZeRO-3 的实现

# ZeRO-3 配置zero_config={"zero_optimization":{"stage":3,# ZeRO-3"reduce_bucket_size":5e8,"allgather_bucket_size":5e8,"contiguous_gradients":True,"overlap_comm":True,"stage3_max_live_parameters":1e9,"stage3_prefetch_bucket_size":5e8,"stage3_param_persistence_threshold":1e6}}

4.5 ZeRO 三阶段对比

阶段参数分片梯度分片优化器分片显存节省通信量
ZeRO-14x2 × Model
ZeRO-28x2 × Model
ZeRO-3Nx3 × Model

5. ZeRO-Offload 与卸载

5.1 ZeRO-Offload 的原理

ZeRO-Offload 将部分计算和存储卸载到 CPU 内存,进一步减少 GPU 显存占用。

5.2 卸载策略

卸载内容卸载到显存节省速度影响
优化器状态CPU减少 50% GPU 显存慢 10-20%
参数CPU减少 33% GPU 显存慢 20-30%
梯度CPU减少 33% GPU 显存慢 20-30%

5.3 ZeRO-Offload 的实现

# ZeRO-3 + Offload 配置zero_config={"zero_optimization":{"stage":3,"offload_optimizer":{"device":"cpu",# 优化器卸载到 CPU"pin_memory":True},"offload_param":{"device":"cpu",# 参数卸载到 CPU"pin_memory":True}}}

5.4 ZeRO-Infinity

ZeRO-Infinity 将卸载扩展到 NVMe 存储,支持千亿参数模型的训练:

存储层级容量带宽延迟存储内容
GPU 显存80GB2 TB/s纳秒当前活跃参数
CPU 内存1TB100 GB/s微秒预取参数
NVMe 存储10TB10 GB/s毫秒不活跃参数

6. DeepSpeed ZeRO 的边界与失效模式

6.1 通信瓶颈

问题表现解决方案
通信量大训练速度慢增加 GPU 数量
通信延迟高同步等待时间长使用更高速网络
通信不平衡某些 GPU 负载高优化通信拓扑

6.2 卸载瓶颈

问题表现解决方案
CPU 带宽不足卸载等待时间长使用更高速 CPU 内存
CPU 内存不足卸载失败增加 CPU 内存
NVMe 带宽不足卸载速度慢使用 NVMe RAID

6.3 DeepSpeed ZeRO 的优缺点总结

优点缺点
显存节省显著通信量增加
支持超大模型实现复杂度高
灵活的分阶段选择小模型收益有限
支持卸载到 CPU/NVMe卸载速度慢

7. DeepSpeed ZeRO 的工程实践

7.1 ZeRO 阶段选择指南

模型规模推荐阶段原因
<1BDDP(ZeRO-0)显存足够,通信少
1B-10BZeRO-2梯度分片,节省显存
10B-100BZeRO-3全参数分片
>100BZeRO-3 + Offload卸载到 CPU/NVMe

7.2 性能优化

优化策略描述效果
通信重叠通信与计算重叠减少 20% 训练时间
梯度累积模拟大 batch提高 GPU 利用率
混合精度BF16 训练减少 50% 显存
参数预取预取下一个模块的参数减少通信等待

7.3 监控与调试

指标描述告警阈值
通信时间通信占总时间比例>30%
显存使用各 GPU 显存使用率>90%
卸载速度CPU/NVMe 卸载速度低于预期 50%

8. ZeRO 的通信模式详解

8.1 ZeRO-1 通信

ZeRO-1 只在优化器更新时需要通信:

defzero1_communication(model,world_size,rank):"""ZeRO-1 通信模式"""# 前向传播:无需通信loss=model.forward(batch)# 反向传播:All-Reduce 梯度(与 DDP 相同)model.backward()# 优化器更新:只更新本分片shard_size=len(model.parameters())//world_size param_shard=list(model.parameters())[rank*shard_size:(rank+1)*shard_size]optimizer.step(param_shard)# 只更新本分片# 收集完整参数forparaminmodel.parameters():dist.all_gather(param,param)
8.2 ZeRO-2 通信

ZeRO-2 在反向传播时使用 Reduce-Scatter 分发梯度:

defzero2_communication(model,world_size,rank):"""ZeRO-2 通信模式"""# 前向传播:无需通信loss=model.forward(batch)# 反向传播:Reduce-Scatter 梯度forparaminmodel.parameters():# 计算梯度后 Reduce-Scattershard_size=param.numel()//world_size chunks=param.grad.view(world_size,shard_size)reduce_scatter_output=torch.zeros(shard_size,device=param.device)dist.reduce_scatter(reduce_scatter_output,chunks)param.grad=reduce_scatter_output# 只保留本分片梯度# 优化器更新:只更新本分片optimizer.step()# 收集完整参数forparaminmodel.parameters():shard_size=param.numel()//world_size shard=param.data[:shard_size]dist.all_gather(param.data.view(world_size,shard_size),shard)
8.3 ZeRO-3 通信

ZeRO-3 在前向和反向传播时都需要 All-Gather:

defzero3_communication(layer,input_data,world_size,rank):"""ZeRO-3 通信模式"""# 前向传播:先收集完整参数shard_size=layer.weight.numel()//world_size shard=layer.weight.data[:shard_size]full_weight=torch.zeros_like(layer.weight.data)dist.all_gather(full_weight.view(world_size,shard_size),shard)# 使用完整参数计算output=layer.forward(input_data)# 丢弃非本分片参数layer.weight.data=shardreturnoutput

9. ZeRO 的卸载策略

9.1 优化器卸载

优化器卸载将 Adam 动量和方差从 GPU 卸载到 CPU 内存:

# ZeRO-Offload 优化器卸载配置zero_config={"zero_optimization":{"stage":3,"offload_optimizer":{"device":"cpu","pin_memory":True,"buffer_count":4,"fast_init":False}}}
卸载策略GPU 显存节省训练速度影响适用场景
无卸载0%基准显存充足
优化器卸载50%慢 10-20%显存不足
优化器+参数卸载66%慢 20-30%显存严重不足
全卸载80%慢 30-50%超大模型
9.2 CPU 优化器计算
defcpu_adam_step(parameters,gradients,optimizer_state):"""CPU 上的 Adam 优化器步骤"""forparam,gradinzip(parameters,gradients):# 在 CPU 上更新参数param.data=param.data-lr*grad/(torch.sqrt(optimizer_state["variance"][param])+1e-8)
9.3 卸载的性能权衡
GPU 显存(GB)可训练模型(ZeRO-3)可训练模型(ZeRO-3 + Offload)
16GB7B13B
32GB13B30B
80GB30B70B
160GB70B175B

10. ZeRO 的训练实践

10.1 训练脚本
importdeepspeeddeftrain_with_deepspeed(model,dataloader,config):"""使用 DeepSpeed ZeRO 训练"""# 初始化 DeepSpeedmodel_engine,optimizer,_,_=deepspeed.initialize(model=model,model_parameters=model.parameters(),config_params=config)forepochinrange(10):forbatchindataloader:loss=model_engine(batch)model_engine.backward(loss)model_engine.step()returnmodel_engine
10.2 配置示例
{"train_batch_size":32,"gradient_accumulation_steps":4,"optimizer":{"type":"AdamW","params":{"lr":1e-4,"weight_decay":0.01}},"zero_optimization":{"stage":3,"offload_optimizer":{"device":"cpu"}},"fp16":{"enabled":true}}
10.3 性能调优
参数推荐值说明
reduce_bucket_size5e8梯度通信 bucket 大小
allgather_bucket_size5e8参数收集 bucket 大小
stage3_prefetch_bucket_size5e8预取 bucket 大小
stage3_max_live_parameters1e9最大存活参数数
gradient_accumulation_steps4梯度积累步数

11. DeepSpeed ZeRO 的进阶功能

11.1 梯度裁剪
# 启用梯度裁剪zero_config={"zero_optimization":{"stage":3,"gradient_clipping":1.0# 梯度裁剪阈值}}
11.2 学习率调度
# 学习率调度配置zero_config={"scheduler":{"type":"WarmupLR","params":{"warmup_min_lr":0,"warmup_max_lr":1e-4,"warmup_num_steps":1000}}}
11.3 混合精度训练
# 混合精度配置zero_config={"bf16":{"enabled":True# 使用 BF16 替代 FP16},"fp16":{"enabled":False}}

总结

DeepSpeed ZeRO 通过分阶段消除数据并行中的冗余存储,将显存占用降低到原来的 1/N。ZeRO-1 分片优化器状态,节省 4x 显存;ZeRO-2 分片梯度,节省 8x 显存;ZeRO-3 全参数分片,节省 Nx 显存。ZeRO-Offload 将计算和存储卸载到 CPU/NVMe,进一步减少 GPU 显存占用。ZeRO 阶段的选择取决于模型规模和硬件资源。

外部引用

  • ZeRO 原始论文:https://arxiv.org/abs/1910.02054
  • DeepSpeed 官方文档:https://www.deepspeed.ai/
  • ZeRO-Offload 卸载:https://arxiv.org/abs/2101.06840
  • ZeRO-Infinity 超大模型:https://arxiv.org/abs/2204.12047
  • DeepSpeed 混合精度:https://www.deepspeed.ai/
  • ZeRO 与 FSDP 对比:https://www.deepspeed.ai/
  • ZeRO-1 优化器分片:https://arxiv.org/abs/1910.02054
  • ZeRO-2 梯度分片:https://arxiv.org/abs/1910.02054
  • ZeRO-3 全参数分片:https://arxiv.org/abs/1910.02054
  • 分布式训练显存优化:https://arxiv.org/abs/2303.04226
http://www.jsqmd.com/news/1332185/

相关文章:

  • 2026 甄选:装配电工 / PLC 编程 / 工业机器人技能培训,长三角五大智能制造实训机构实力深度剖析 - 甄选测评馆
  • 2026年沈阳防雷检测机构挑选攻略 中科智电等合规企业盘点 - 资讯在线
  • 终极iOS虚拟定位工具:iFakeLocation跨平台使用完整指南
  • OceanBase 可用区和节点的管理
  • 从“云原生“到“AI原生“:2026年云原生架构的范式跃迁与工程实践
  • 2024 Python爬虫实战指南:从基础到工程化,应对反爬与动态渲染
  • 类型安全都一样,单调用却慢 41 倍:Agent 工具调用 JSON 校验的实测复盘
  • 2026上海宝格丽回收避坑指南:31年金字招牌易奢福,为交易保驾护航 - 奢侈品回收探店ing
  • 【客户定制更新】智慧城市运行管理服务平台版本更新内容——全局优化、业务模块优化
  • 从黑盒到白盒:逆向分析赛尔号通信协议的技术实践
  • 百度网盘下载加速终极教程:告别限速,5步获取真实下载地址
  • 用SDD与Spec-Kit驯服AI编码幻觉:从模糊需求到精准代码生成
  • 图像融合算法全解析:从像素级到决策级,实战指南与避坑心得
  • 2026年外贸拓客软件避坑全攻略:跨境魔方领衔区分正规海关数据平台与劣质线索工具 警惕虚假邮箱与高额年费陷阱
  • OpenIM Server v3.8.3-patch.16深度解析:性能优化与稳定性加固实战
  • 逆向解析微信读书API:从抓包到实现个人数据同步与自动化
  • Prompt Engineer实战指南:从原理到应用,掌握与大模型高效沟通的核心方法
  • 2026许昌设计能力强的防碱防潮浓缩液定制厂家推荐 - 汇聚至此
  • 从OpenClaw到LightVela:AI Agent开发的可视化配置与效率提升实践
  • 淘宝改价系统:批量改价3秒完成1000品,竞品没反应过来你就调完了
  • ATV900变频器在起重行业的抱闸控制与安全应用
  • 第8讲:MCP 协议——给 Agent 接上真实系统
  • 地产沙盘定制真实口碑:本地服务商项目落地体验一览 - 优企甄选
  • AUTOSAR DEM模块Operation Cycle:诊断事件状态管理与老化机制详解
  • HAProxy负载均衡核心配置与性能优化实战
  • SQL注入从原理到实战:基于DVWA靶场的漏洞剖析与防御指南
  • Windows To Go实战指南:打造便携式Windows系统盘,实现跨设备无缝工作
  • 从AI代码生成到工程化交付:构建可控的AI编程工作流
  • BRFSS数据集解析:公共卫生数据分析与应用
  • 2026年跨境魔方B2B外贸拓客工具横评:海关数据社媒谷歌搜索合规选型指南