· 分布式训练

ZeRO-3 核心流程:把参数也切开的那一档

参数、梯度、优化器状态全部按卡数切开,每卡只留 1/N;哪一层要用就临时把参数拼回来,算完立刻释放。

一句话:把"参数、梯度、优化器状态"三样东西全部按卡数 N 切开,每卡只留 1/N;哪一层要用,就临时把参数拼回来,算完立刻扔掉。

1 到底切开了什么

DDP 与 ZeRO 三阶段分别切开了什么DDP 什么都不切,ZeRO-1 切优化器状态,ZeRO-2 再切梯度,ZeRO-3 连参数也切方案参数梯度优化器状态每卡常驻显存基线(DDP)完整一份完整一份完整一份120 GBZeRO-1完整一份完整一份31.4 GBZeRO-2完整一份16.6 GBZeRO-31.9 GB灰条 = 每张卡都完整存一份;8 个小格 = 按 N=8 切开,深色格是这张卡自己留的那 1/N显存数字 = 7.5B 模型 + Adam fp32 优化器状态、64 卡(DeepSpeed 论文实测值)

显存最大的一块其实不是参数,而是优化器状态——Adam 的 fp32 主权重 + 动量 + 方差,每个参数要 12 字节,比参数本身(fp16,2 字节)重 6 倍。所以 ZeRO 是先从最肥的地方切起。

展开:120GB 是怎么算出来的,以及显存随卡数怎么降
项目精度7.5B 模型占用
参数fp1615 GB
梯度fp1615 GB
优化器状态(主权重 + 动量 + 方差)fp3290 GB
合计120 GB
DDP(不切) ZeRO-1 ZeRO-2 ZeRO-3
0.251416642568163264128256512
双对数坐标。ZeRO-3 是一条斜率为 −1 的直线(120/N),卡数翻倍、显存减半;ZeRO-1/2 到了后面会被"没被切的那部分"拖住,降不下去。

2 每层怎么做

ZeRO-3 每层要走的四个动作前向拼参数算完即扔,反向再拼一次并分发梯度,优化器只用本地分片更新前向 · 拼参数All-Gather 该层拼出完整一层算完 · 立刻扔只留自己那 1/N显存就省在这反向 · 再拼一次参数已被扔掉算完梯度就分发更新 · 各改各的只用本地 1/N 分片这一步零通信
  1. 前向:All-Gather 该层的参数分片,临时拼出完整一层 → 算完立刻释放,只留自己那 1/N。
  2. 反向:参数上一步已经扔了,所以再 All-Gather 一次 → 算梯度 → 立刻 Reduce-Scatter,每卡只留 1/N。
  3. 优化器:只用本地 1/N 参数片 + 本地优化器状态更新,零通信
  4. 刚更新的分片,正好是下一轮前向要拼的那一份 —— 闭环。

3 多出来的通信藏在哪

参数 All-Gather 与前向计算并行算第 i 层时提前拼第 i+1 层的参数,通信被计算完全盖住前向计算参数 All-Gather(提前发起)算第 1 层算第 2 层算第 3 层拼第 2 层参数拼第 3 层参数拼第 4 层参数上下两块同宽 = 完全并行。只有第一层和最后一次反向没有计算可挡,那点延迟会露在外面
算第 i 层时就提前发起第 i+1 层的 All-Gather(prefetch),上下两块完全并行。
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 官方的实现是 FSDPShardingStrategy.FULL_SHARD);组内 FSDP + 组间 DDP 的混合模式叫 HSDPHYBRID_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 桶大小,最后才动其他开关。粒度不对的话,后面怎么调都救不回来。

分享:
返回文章列表

相关文章

全部文章 »