· 分布式训练
ZeRO-3 核心流程:把参数也切开的那一档
参数、梯度、优化器状态全部按卡数切开,每卡只留 1/N;哪一层要用就临时把参数拼回来,算完立刻释放。
一句话:把"参数、梯度、优化器状态"三样东西全部按卡数 N 切开,每卡只留 1/N;哪一层要用,就临时把参数拼回来,算完立刻扔掉。
1 到底切开了什么
显存最大的一块其实不是参数,而是优化器状态——Adam 的 fp32 主权重 + 动量 + 方差,每个参数要 12 字节,比参数本身(fp16,2 字节)重 6 倍。所以 ZeRO 是先从最肥的地方切起。
展开:120GB 是怎么算出来的,以及显存随卡数怎么降
| 项目 | 精度 | 7.5B 模型占用 |
|---|---|---|
| 参数 | fp16 | 15 GB |
| 梯度 | fp16 | 15 GB |
| 优化器状态(主权重 + 动量 + 方差) | fp32 | 90 GB |
| 合计 | 120 GB |
DDP(不切) ZeRO-1 ZeRO-2 ZeRO-3
2 每层怎么做
- 前向:All-Gather 该层的参数分片,临时拼出完整一层 → 算完立刻释放,只留自己那 1/N。
- 反向:参数上一步已经扔了,所以再 All-Gather 一次 → 算梯度 → 立刻 Reduce-Scatter,每卡只留 1/N。
- 优化器:只用本地 1/N 参数片 + 本地优化器状态更新,零通信。
- 刚更新的分片,正好是下一轮前向要拼的那一份 —— 闭环。
3 多出来的通信藏在哪
1/N每卡常驻显存
3 : 2通信量 vs DDP = 1.5×
≈0预取后暴露在外的增量
DDP 只有 1 次 AllReduce(2 个单位),ZeRO-3 是「前向 AG + 反向 AG + 反向 RS」共 3 个单位。多出来的就是"反向再拼一次参数"——因为前向算完就把参数扔了。靠预取基本能藏掉,只有第一层和最后一次反向没有计算可挡。
4 和邻居的关系
| 方案 | 切什么 | 一句话定位 |
|---|---|---|
| DDP | 都不切 | 每卡一份完整副本,通信最省、显存最费 |
| ZeRO-1 | 优化器状态 | 性价比最高的一档,通信量和 DDP 一样 |
| ZeRO-2 | + 梯度 | 通信量仍等于 DDP,显存再降一截 |
| ZeRO-3 / FSDP | + 参数 | 显存降到 1/N,代价是通信变成 1.5× |
PyTorch 官方的实现是 FSDP(ShardingStrategy.FULL_SHARD);组内 FSDP + 组间 DDP 的混合模式叫 HSDP(HYBRID_SHARD),多机时比纯 FSDP 省一半跨机通信。
坑 1:梯度裁剪必须分片感知。各卡只看到自己 1/N 的范数,直接
clip_grad_norm_ 会让每张卡用不同的裁剪系数,权重静默发散——不报错,只是 loss 悄悄变差。坑 2:通信粒度不能太碎。前向算完扔参数意味着反向要重算,同时会产生大量"层粒度"的小 All-Gather。层切得太细或 bucket 太小,通信效率直接掉下来。
坑 3:和前向重算的叠加。用了 activation checkpointing 后,重算那一段还要再走一次参数 All-Gather,显存省了但通信又多一截,两个一起开需要单独压测。
展开:FSDP 与 DeepSpeed 的关键开关
# PyTorch FSDP (= ZeRO-3)
FSDP(
model,
sharding_strategy=ShardingStrategy.FULL_SHARD, # 纯 ZeRO-3;HYBRID_SHARD = HSDP
auto_wrap_policy=transformer_auto_wrap_policy( # 按 Transformer block 分片
transformer_layer_cls={LlamaDecoderLayer}),
backward_prefetch=BackwardPrefetch.BACKWARD_PRE, # 反向预取,藏通信的关键
limit_all_gathers=True, # 控制显存峰值
use_orig_params=True, # 允许对参数做特殊处理(如 LoRA)
)// DeepSpeed ZeRO-3
{
"zero_optimization": {
"stage": 3,
"stage3_prefetch_bucket_size": 5e7, // 预取桶:越大越省延迟,越费显存
"stage3_param_persistence_threshold": 1e5, // 太小的参数不切,避免碎片化通信
"stage3_max_live_parameters": 1e9, // 同时在显存里的完整参数上限
"reduce_bucket_size": 5e6,
"contiguous_gradients": true
}
}调参优先级:先把 auto_wrap 的粒度调对(一层一个 FSDP unit),再看 prefetch 桶大小,最后才动其他开关。粒度不对的话,后面怎么调都救不回来。