AI ZeRO优化器实践笔记
ZeRO(Zero Redundancy Optimizer)是DeepSpeed里的一组优化策略,核心思想就是把原本每个GPU都要存的东西拆散,按需分发给不同GPU。
我一开始用的是PyTorch原生的DistributedDataParallel(DDP),显存确实不够。
ZeRO优化器是什么
ZeRO(Zero Redundancy Optimizer)是DeepSpeed里的一组优化策略,核心思想就是把原本每个GPU都要存的东西拆散,按需分发给不同GPU。
传统数据并行是这样的:
ZeRO把这些冗余去掉了:
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_optimizer和offload_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_steps和train_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本身会加剧这个问题,因为它频繁地分配和释放内存块。我的解决办法是:
- 关闭不必要的显存缓存:
torch.cuda.empty_cache()
调整
memory_efficient_attention参数(如果用的是transformer模型)尝试不同的
bucket_size设置,这个参数会影响通信时的内存分配
实际效果对比
以7B模型为例,我对比了几种配置:
| 配置 | 每GPU显存占用 | 每步时间 | 有效batch size |
|---|---|---|---|
| DDP | 38GB | 120ms | 32 |
| ZeRO Stage 1 | 28GB | 125ms | 64 |
| ZeRO Stage 2 | 22GB | 130ms | 128 |
| ZeRO Stage 2 + offload | 18GB | 180ms | 128 |
| ZeRO Stage 3 | 15GB | 210ms | 256 |
显存占用确实下降明显,但训练时间也相应增加了。这里有个trade-off需要权衡。
我的结论是:对于大多数情况,ZeRO Stage 2是个不错的选择。如果显存真的不够,再考虑offload。Stage 3除非有特殊需求,否则不太建议。
一些小技巧
- 监控显存使用:
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")
逐步增加batch size:不要上来就设置大的batch size,先从小开始,确认能跑起来了再逐步增加。
检查模型是否正确切分:可以打印模型参数,看是否正确分布在不同GPU上:
for name, param in model.named_parameters():
print(f"{name}: {param.device}, {param.shape}")
- 使用
torch.profiler:分析训练中的瓶颈,看是不是通信占了太多时间。
最后一点建议
ZeRO是个很实用的工具,但它不是万能的。有时候问题出在别的地方,比如模型结构本身就不高效,或者数据预处理有问题。
我见过有人用ZeRO强行跑一个超大的模型,结果训练出来效果很差。后来发现是因为模型太大了,训练数据根本不够撑住这么大的参数量。
所以工具得用在合适的地方。ZeRO解决了显存问题,但能不能用好,还得看你的模型和数据。
这大概是我在训练里学到的最重要的一点:工具很重要,但更重要的是知道什么时候用什么工具,以及什么时候不用什么工具。
版权声明: 本文首发于 指尖魔法屋-AI ZeRO优化器实践笔记(https://blog.thinkmoon.cn/post/999-ai-zero-optimizer-training-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。