1. 从“年糕切片”理解ZeRO的核心思想

想象你面前有一块完整的年糕,要分给四个朋友吃。最直接的方法是每人切一块,但这样每人手里都拿着完整的年糕块,既占空间又浪费资源。ZeRO(Zero Redundancy Optimizer)的核心理念就像把年糕切成薄片分发:每张显卡只保存模型的一部分状态,需要时再临时拼合。

具体来说,大模型训练时会占用显存的四大金刚是:

  • 模型参数(W):FP16精度下6B参数约12GB
  • 梯度(G):反向传播产生的中间结果,同样占12GB
  • 优化器状态(OS):Adam等优化器需要保存的动量等参数,FP32精度下约24GB
  • 激活值(Activation):前向传播的中间结果,随序列长度指数增长

传统数据并行就像人手一份完整年糕,每张显卡都要保存全部W+G+OS。而ZeRO通过三个阶段的分片策略,让显存占用从O(N)降到O(1/N):

# 传统数据并行的显存占用
total_memory = (W + G + OS) * num_gpus  

# ZeRO分片后的显存占用 
zero_memory = W/num_gpus + G/num_gpus + OS/num_gpus

2. ZeRO三阶段的切片逻辑详解

2.1 ZeRO-1:优化器状态分片

就像把年糕的调味料分开保存,ZeRO-1只对优化器状态(如Adam的m/v矩阵)做分片。每张显卡只需维护自己那部分参数的优化状态,显存节省约4倍。实际使用时你会发现:

  • 通信开销:仅需在更新参数时同步梯度(AllReduce)
  • 适用场景:当你的显存主要被优化器占用时(比如使用AdamW+FP32)

2.2 ZeRO-2:梯度分片进阶

在ZeRO-1基础上,进一步把梯度也切片保存。反向传播时,每张卡只计算自己负责的那部分参数的梯度。这里有个关键细节:

# deepspeed.json配置片段
"zero_optimization": {
  "stage": 2,
  "reduce_bucket_size": 5e8  # 控制通信时缓冲区大小
}
  • reduce_bucket_size参数决定了每次通信传输的数据量,太大会占内存,太小会增加通信次数
  • 实测在A100上设置为500MB时,训练速度比默认值快15%

2.3 ZeRO-3:参数分片终极版

最彻底的切片方案,把模型参数也分散到各卡。需要前向计算时,通过AllGather操作临时重构完整参数。这就像:

  1. 每人只保管年糕的某几层切片
  2. 需要吃的时候把所有人的切片叠起来
  3. 吃完再把各自那层收好

典型的内存-通信权衡案例:

  • 显存占用降至1/num_gpus
  • 但每个step要多出2次AllGather通信
  • 建议在40B以上模型使用

3. Offload:当切片还不够时的终极武器

有时候即使切成极薄片,显存还是装不下。这时就该祭出Offload大招——把切片好的年糕暂时放到冰箱(CPU内存)里保存。Deepspeed支持两种卸载方式:

卸载类型保存位置典型场景速度影响
OptimizerCPU优化器状态太大20%↓
ParameterCPU/NVMe百亿参数模型50%↓
混合卸载分层存储超大模型+有限硬件30-70%↓

配置示例:

{
  "zero_optimization": {
    "stage": 3,
    "offload_optimizer": {
      "device": "cpu",
      "buffer_count": 4  // 控制异步卸载的缓冲区数量 
    }
  }
}

实测在LLaMA-7B微调时,ZeRO-3+Offload可以让单卡显存从48GB降到12GB,但训练速度会降低约40%。

4. ChatGLM2-6B微调实战演示

4.1 环境准备

建议使用AutoDL云服务(没有广告费纯推荐):

# 基础环境
conda create -n chatglm python=3.9
pip install torch==1.12.1+cu116 --extra-index-url https://download.pytorch.org/whl/cu116
git clone https://github.com/THUDM/ChatGLM2-6B
cd ChatGLM2-6B
pip install -r requirements.txt

4.2 关键配置解析

修改ds_train_finetune.sh时注意:

--per_device_train_batch_size 4  # 根据显存动态调整
--gradient_accumulation_steps 2  # 模拟更大batch size
--fp16  # 必须开启以配合ZeRO

deepspeed.json的黄金参数组合:

{
  "zero_optimization": {
    "stage": 2,
    "offload_optimizer": {
      "device": "cpu"
    },
    "allgather_bucket_size": 2e8,
    "reduce_bucket_size": 2e8
  },
  "gradient_clipping": 1.0,
  "train_micro_batch_size_per_gpu": "auto"
}

4.3 避坑指南

  1. 遇到"CUDA out of memory"时:
    • 先尝试减小batch_size
    • 然后考虑开启ZeRO-2
    • 最后才用Offload
  2. 训练速度异常慢检查:
    • nccl版本是否匹配
    • 通信带宽是否被占满
  3. 精度问题:
    • FP16训练可能需要调loss scaling
    • 尝试--bf16如果硬件支持

5. 性能调优经验谈

在8张A100上实测ChatGLM2-6B微调:

  • ZeRO-1:显存18GB/卡,吞吐量120 samples/sec
  • ZeRO-2:显存12GB/卡,吞吐量95 samples/sec
  • ZeRO-3+Offload:显存6GB/卡,吞吐量40 samples/sec

推荐的选择策略:

  1. 显存充足优先用ZeRO-1
  2. batch_size需要>8时用ZeRO-2
  3. 只有少量高端卡时考虑ZeRO-3
  4. 消费级显卡(如3090)必须Offload

最后分享一个监控显存的神器:

nvidia-smi -l 1  # 实时查看显存变化
watch -n 1 'gpustat -cup'  # 更直观的显示

更多推荐