图解Deepspeed ZeRO:从“年糕切片”到实战调优,轻松搞定大模型训练
·
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/num_gpus
- 但每个step要多出2次AllGather通信
- 建议在40B以上模型使用
3. Offload:当切片还不够时的终极武器
有时候即使切成极薄片,显存还是装不下。这时就该祭出Offload大招——把切片好的年糕暂时放到冰箱(CPU内存)里保存。Deepspeed支持两种卸载方式:
| 卸载类型 | 保存位置 | 典型场景 | 速度影响 |
|---|---|---|---|
| Optimizer | CPU | 优化器状态太大 | 20%↓ |
| Parameter | CPU/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 避坑指南
- 遇到"CUDA out of memory"时:
- 先尝试减小batch_size
- 然后考虑开启ZeRO-2
- 最后才用Offload
- 训练速度异常慢检查:
- nccl版本是否匹配
- 通信带宽是否被占满
- 精度问题:
- 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
推荐的选择策略:
- 显存充足优先用ZeRO-1
- batch_size需要>8时用ZeRO-2
- 只有少量高端卡时考虑ZeRO-3
- 消费级显卡(如3090)必须Offload
最后分享一个监控显存的神器:
nvidia-smi -l 1 # 实时查看显存变化
watch -n 1 'gpustat -cup' # 更直观的显示
更多推荐
所有评论(0)