AI ZeRO优化器实践笔记

ZeRO(Zero Redundancy Optimizer)是DeepSpeed里的一组优化策略,核心思想就是把原本每个GPU都要存的东西拆散,按需分发给不同GPU。

我一开始用的是PyTorch原生的DistributedDataParallel(DDP),显存确实不够。

ZeRO优化器是什么

ZeRO(Zero Redundancy Optimizer)是DeepSpeed里的一组优化策略,核心思想就是把原本每个GPU都要存的东西拆散,按需分发给不同GPU。

传统数据并行是这样的:

graph LR A[GPU 0] -->|完整模型+优化器+梯度| B[独立计算] C[GPU 1] -->|完整模型+优化器+梯度| D[独立计算] E[GPU 2] -->|完整模型+优化器+梯度| F[独立计算] G[GPU 3] -->|完整模型+优化器+梯度| H[独立计算]

ZeRO把这些冗余去掉了:

graph LR A[GPU 0<br/>优化器状态0] -->|需要时交换| B[GPU 1<br/>优化器状态1] B -->|需要时交换| C[GPU 2<br/>优化器状态2] C -->|需要时交换| D[GPU 3<br/>优化器状态3]

ZeRO有3个stage,逐级递进:

  • Stage 1:只拆分优化器状态(Adam的动量和方差)
  • Stage 2:在Stage 1基础上再拆分梯度
  • Stage 3:在前两个基础上再拆分模型参数

理论上能节省的显存比例大约是:Stage 1省4倍,Stage 2省8倍,Stage 3省更多,但通信开销也更大。

实践中的配置

我一开始用的是PyTorch原生的DistributedDataParallel(DDP),显存确实不够。后来换成DeepSpeed ZeRO Stage 2,情况就完全不同了。

配置文件大概是这样的:

{
  "train_batch_size": 128,
  "train_micro_batch_size_per_gpu": 4,
  "gradient_accumulation_steps": 8,
  "optimizer": {
    "type": "AdamW",
    "params": {
      "lr": "1e-4",
      "betas": [0.9, 0.999],
      "eps": "1e-8",
      "weight_decay": "0.01"
    }
  },
  "scheduler": {
    "type": "WarmupLR",
    "params": {
      "warmup_min_lr": "0",
      "warmup_max_lr": "1e-4",
      "warmup_num_steps": 10000
    }
  },
  "zero_optimization": {
    "stage": 2,
    "offload_optimizer": {
      "device": "cpu"
    },
    "offload_param": {
      "device": "cpu"
    }
  },
  "fp16": {
    "enabled": true,
    "loss_scale": 0,
    "initial_scale_power": 16,
    "loss_scale_window": 1000,
    "hysteresis": 2,
    "min_loss_scale": 1
  },
  "activation_checkpointing": {
    "partition_activations": true,
    "cpu_checkpointing": true,
    "contiguous_memory_optimization": true
  }
}

这里有个关键配置:offload_optimizeroffload_param。我把优化器状态和模型参数都offload到CPU了,虽然慢了一点,但显存确实省了不少。

Stage选择的经验

很多人上来就问该用哪个stage。这个问题其实没有标准答案,得看具体情况。

我一开始图省事直接上了Stage 3,结果训练速度比Stage 2慢了快一倍。后来查了监控才发现,Stage 3在每一步都要在GPU之间同步模型参数,通信开销实在太大。

所以我的经验是:

先用Stage 2,显存不够再加offload。

只有在以下情况才考虑Stage 3:

  • 模型确实很大,Stage 2+offload都跑不起来
  • 通信带宽很充足(比如NVLink)
  • 你有足够的GPU数量,通信成本相对可控

另外,offload也要慎用。我试过把所有东西都offload到CPU,结果训练速度慢得离谱。后来只offload优化器状态,速度就回到了可接受范围。

踩过的坑

梯度累积计算错误

一开始配置ZeRO的时候,我把gradient_accumulation_stepstrain_batch_size搞混了。本来应该是128/8=16,结果写成了128/4=32。

这个问题很难察觉,因为训练能跑起来,但效果就是不对。后来写了个简单的脚本验证:

import torch.distributed as dist

# 检查batch size是否正确
if dist.get_rank() == 0:
    print(f"Total batch size: {args.train_batch_size}")
    print(f"Micro batch size: {args.train_micro_batch_size_per_gpu}")
    print(f"Gradient accumulation: {args.gradient_accumulation_steps}")
    print(f"World size: {dist.get_world_size()}")
    print(f"Effective batch size: {args.train_micro_batch_size_per_gpu * args.gradient_accumulation_steps * dist.get_world_size()}")

FP16溢出问题

ZeRO配合FP16的时候,经常会遇到loss scaling的问题。我一开始用的动态loss scaling,结果训练到一半loss就变成NaN了。

后来改成固定loss scaling:

"fp16": {
  "enabled": true,
  "loss_scale": 65536,
  "initial_scale_power": 16,
  "loss_scale_window": 1000,
  "hysteresis": 2,
  "min_loss_scale": 1
}

但固定scaling也不是万能的,如果还是溢出,就得把learning rate调小一些,或者检查数据里有没有异常值。

显存碎片化

有时候总显存看起来够用,但还是爆显存了。这多半是显存碎片化的问题。

ZeRO本身会加剧这个问题,因为它频繁地分配和释放内存块。我的解决办法是:

  1. 关闭不必要的显存缓存:
torch.cuda.empty_cache()
  1. 调整memory_efficient_attention参数(如果用的是transformer模型)

  2. 尝试不同的bucket_size设置,这个参数会影响通信时的内存分配

实际效果对比

以7B模型为例,我对比了几种配置:

配置每GPU显存占用每步时间有效batch size
DDP38GB120ms32
ZeRO Stage 128GB125ms64
ZeRO Stage 222GB130ms128
ZeRO Stage 2 + offload18GB180ms128
ZeRO Stage 315GB210ms256

显存占用确实下降明显,但训练时间也相应增加了。这里有个trade-off需要权衡。

我的结论是:对于大多数情况,ZeRO Stage 2是个不错的选择。如果显存真的不够,再考虑offload。Stage 3除非有特殊需求,否则不太建议。

一些小技巧

  1. 监控显存使用
import torch

def print_memory_stats():
    allocated = torch.cuda.memory_allocated() / 1024**3
    reserved = torch.cuda.memory_reserved() / 1024**3
    print(f"Allocated: {allocated:.2f} GB, Reserved: {reserved:.2f} GB")
  1. 逐步增加batch size:不要上来就设置大的batch size,先从小开始,确认能跑起来了再逐步增加。

  2. 检查模型是否正确切分:可以打印模型参数,看是否正确分布在不同GPU上:

for name, param in model.named_parameters():
    print(f"{name}: {param.device}, {param.shape}")
  1. 使用torch.profiler:分析训练中的瓶颈,看是不是通信占了太多时间。

最后一点建议

ZeRO是个很实用的工具,但它不是万能的。有时候问题出在别的地方,比如模型结构本身就不高效,或者数据预处理有问题。

我见过有人用ZeRO强行跑一个超大的模型,结果训练出来效果很差。后来发现是因为模型太大了,训练数据根本不够撑住这么大的参数量。

所以工具得用在合适的地方。ZeRO解决了显存问题,但能不能用好,还得看你的模型和数据。

这大概是我在训练里学到的最重要的一点:工具很重要,但更重要的是知道什么时候用什么工具,以及什么时候不用什么工具。

版权声明: 本文首发于 指尖魔法屋-AI ZeRO优化器实践笔记https://blog.thinkmoon.cn/post/999-ai-zero-optimizer-training-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!