ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

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

【硬核拆解】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
返回列表