首页
/ DeepSpeed ZeRO 通信优化深度解析:用 reduce-scatter 分区感知规约替代 all-reduce

DeepSpeed ZeRO 通信优化深度解析:用 reduce-scatter 分区感知规约替代 all-reduce

2026-09-08 20:09:12作者:滑思眉Philip

本文基于仓库内发布于 2020 年 3 月的公告 docs/_posts/2020-03-17-reduce-scatter.md 展开。该公告预告了 DeepSpeed ZeRO 第一阶段引入"分区感知(partition-aware)"的梯度规约方式,用 reduce-scatter 取代初版实现中"一次性全局 all-reduce"的通信模式,宣称可将总通信量从数据并行基线的 1.5 倍降至 1.0 倍、通信耗时最多降低 2 倍。时至今日,这一设计已成为仓库中 ZeRO Stage 1/2 的默认梯度平均实现。阅读本文后,你将理解为什么梯度平均不需要 all-reduce、reduce-scatter 如何与 ZeRO 的参数分区天然契合,以及当前仓库源码中该路径的实现与相关配置参数。

一、这篇预览公告的背景与三大结论

ZeRO stage 1 with reduced communication 是 DeepSpeed 官方在早期发布的一篇"预览(sneak preview)"性质公告,正文以三个要点概括了 ZeRO 训练通信优化的核心思路:

  • 分区感知(partition-aware)方案取代了初版实现采用的全局集合通信(all-reduce);
  • 总通信量由数据并行(data parallelism)的 1.5 倍降至 1.0 倍
  • 相比 all-reduce,通信时间最多可降低 2 倍

需要说明的是,上述 1.5x、2x 等量化结论是官方发布该公告时的口径与测量结果,具体数值会随模型规模、GPU 数量、梯度桶(bucket)切分策略等条件而变化;本文后续将聚焦"为什么能省通信"的机制,这部分在当前仓库源码中有完整、可直接核验的实现证据。

二、机制背景:为什么梯度平均用 all-reduce 是"浪费"的

在标准数据并行训练中,每个 rank 在各自数据分片上反向传播后,需要把各副本上相同参数的梯度取平均,再用于更新参数。最直接的做法是对整份梯度做一次 all-reduce(Ring All-Reduce)。从通信原语上看,一次 all-reduce 在功能上等价于 reduce-scatter + all-gather:先把各 rank 的数据规约并"打散"到所有 rank(每 rank 得到全量结果的一个分片),再通过 all-gather 把完整结果广播到每个 rank。

ZeRO 的省显存思路是"分区":优化器状态(Stage 1)、乃至 16-bit 梯度(Stage 2)都被切分成 数据并行度 等份,每个 rank 只负责持有并更新自己那份分区。因此,梯度平均的真正产物并不需要是全量平均梯度——每个 rank 只需要"与自己持有的参数分区相对应"的那一段平均梯度切片即可完成本地优化器更新。

于是问题就显现了:all-reduce 多做了最后一步 all-gather,把每个 rank 都不需要的其它分区数据也"全员广播"了一遍。这正是公告所说的"初始实现使用全局集体通信(all-reduce)"存在的通信冗余,而 reduce-scatter 只做"规约 + 打散",恰好把平均梯度的最终落点对准各自的分区,从而省去 all-gather 阶段。

三、当前仓库中的实现佐证:分区感知规约已写入 ZeRO 优化器

该公告预告的优化并非停留在概念上。在当前仓库中,ZeRO Stage 1 与 Stage 2 由同一个优化器类统一实现: deepspeed/runtime/zero/stage_1_and_2.py。从源码可以清晰看到分区感知规约的落地:

1. reduce_scatter 是默认开启的开关

构造参数中 reduce_scatter=True 为默认值(stage_1_and_2.py#L164),并在初始化时被保存为 self.reduce_scatterstage_1_and_2.py#L230)。同一文件中用 partition_gradients 区分两个阶段:True 时为 ZeRO-2(梯度也被分区),False 时为 ZeRO-1(仅优化器状态分区),见 stage_1_and_2.py#L224-L226。可见"分区感知、只保留本地分区所需梯度切片"的逻辑对两个阶段是统一的。

2. average_tensor:一个方法内的两条规约路径

梯度规约的核心函数是 average_tensorstage_1_and_2.py#L1360-L1478):

  • 关闭 reduce_scatter 时,走 gradient_reduction_w_predividestage_1_and_2.py#L1269-L1298),内部对整桶梯度调用 dist.all_reduce,并通过 gradient_predivide_factorpostscale_gradientsgradient_average 等参数在 fp16 下做"先除后归约"以控制数值稳定性;
  • 开启 reduce_scatter 时,代码遍历桶内每个参数的分区元数据 param_to_partition_idsgrad_start_offset,把梯度张量按目标分区 (dst_rank, bucket_offset, numel) 切成若干连续切片(stage_1_and_2.py#L1386-L1442),随后按切片目标分组:
    • 切片只属于单一目标 rank 时,使用 ("reduce", dst, process_group) 为键、对该目标做 dist.reduce(经 allreduce_no_retainallreduce_bucket(rank=dst),其内部即 stage_1_and_2.py#L1866-L1871dist.reduce 到目标全局 rank);
    • 存在多 rank 副本需求等特殊场景时,才回退到 allreduce_and_scatter 路径(stage_1_and_2.py#L1466-L1470)。

换言之,一个参数可能横跨多个分区、一个"规约桶"里也往往有多个连续片段指向同一目标 rank,代码会将这些片段合并后再一次性规约(stage_1_and_2.py#L1435-L1442),尽量让每次集合通信都有足够大的数据量。

3. 归约结果直接对位本地分区,非本地梯度尽早释放

reduce_ipg_gradsstage_1_and_2.py#L1701-L1765)中,规约完成后若 partition_gradients 为真,代码会:

  • 对不属于当前 rank 分区的参数,直接清空其梯度(clear_grad_attribute),从而在不持有完整平均梯度的前提下省下显存;
  • 对属于本地分区的参数,调用 copy_grads_in_partition 把切片写入本地连续分区缓冲,供后续更新该分区的优化器状态使用。

ZeRO-2"只保留自己那份梯度"的显存收益,正是以这套分区感知的规约落位为前提的。文档 docs/_tutorials/zero.md 对 ZeRO 各阶段的划分给出了同一口径的描述:Stage 2 即"规约后的 16-bit 梯度也被分区,每个进程只保留与其优化器状态分区对应的那部分梯度"。

四、与反向传播重叠:IPG 桶与 overlap_comm

reduce-scatter 不止省通信量,还便于与反向传播重叠。ZeRO Stage 1/2 实现了"独立分区梯度桶(IPG,Independent Partition Gradient)"机制:每个参数的反向梯度一经产生,就立即按分区切片拷贝进 reduce_bucket_size 大小的连续桶(stage_1_and_2.py#L1202-L1263),当桶满或本轮 backward 结束时触发 average_tensor 执行上述分区感知规约。开启 overlap_comm 后,规约在独立的 reduction 流(stream)上进行(stage_1_and_2.py#L1360-L1375),从而把梯度通信"藏"在后续反向计算背后,进一步摊薄通信开销。

五、配置参数与启用方式

在 DeepSpeed 配置中,只需为 zero_optimization 启用 ZeRO 即可使用该机制;reduce_scatter 相关项均位于 zero_optimization 键下。以下为参考配置(融合了文档 docs/_tutorials/zero.md 中 Stage 1 与 Stage 2 的示例写法):

{
  "train_batch_size": 32,
  "gradient_accumulation_steps": 1,
  "zero_optimization": {
    "stage": 1,
    "reduce_bucket_size": 5e8,
    "contiguous_gradients": true,
    "reduce_scatter": true,
    "overlap_comm": true
  }
}

其中各字段的含义与默认值(见 docs/_pages/config-json.md#L514-L536deepspeed/runtime/zero/config.py#L108-L113):

参数 说明 默认值
reduce_scatter 是否用 reduce/reduce-scatter 代替 all-reduce 平均梯度(即分区感知规约的开关,对应本公告主题) true
reduce_bucket_size 单次规约处理的元素数上限,限制一次集合通信占用的内存(桶大小) 5e8(约 5 亿元素)
contiguous_gradients 反向过程中把梯度拷贝进连续缓冲,避免内存碎片化 true
overlap_comm 是否尝试将梯度规约与反向计算重叠 false
stage 1 即 ZeRO Stage 1(仅分区优化器状态);取 2 时配合 reduce_scatter: true 同时分区梯度 无(必填)

要点提示:

  • 字段描述与文档口径一致:"Uses reduce or reduce scatter instead of allreduce to average gradients",默认即开启(docs/_pages/config-json.md#L520-L524)。因此在当前版本的 DeepSpeed 中,本文讨论的通信优化是默认生效的梯度平均路径;
  • reduce_bucket_size 是影响通信效率与内存的关键调优项:桶越大,单次集合通信的效率越高,但临时缓冲占用也越大;它对上文"分片后合并、尽量一次规约"的合并效率有直接约束;
  • overlap_comm 需要与 IPG 桶、连续梯度等机制配合才能发挥效果,并非所有后端/加速器都支持流级重叠。

六、正确性保护与扩展场景

分区感知规约在实现上还包含若干正确性与扩展性设计,同样可在源码中核验:

  • 均值语义:开启 reduce_scatter 时,先按 dp_world_size / sequence_parallel_size 对梯度做除法再发往目标分区(stage_1_and_2.py#L1445-L1446),保证结果等价于全体数据并行副本的平均;
  • fp16 下的数值安全:关闭 reduce_scatter 的 all-reduce 路径保留了 gradient_predivide_factorpostscale_gradients 的前置/后置缩放,这是 fp16 大世界规模梯度规约的经典数值保护手段;
  • MoE / 专家并行:当桶内存在 MoE 参数时,会切换到专家数据并行进程组执行规约(stage_1_and_2.py#L1399-L1401),避免在错误的通信域内混算;
  • 序列并行:规约前的除数使用 dp_world_size / sequence_parallel_size,说明该通信路径与序列并行维度按设计协作,且代码在序列并行大于 1 时会把通信 dtype 提升为 fp32(stage_1_and_2.py#L1857-L1858)。

七、总结

从 2020 年 3 月这篇"分区感知、降低通信"的预览公告,到当前仓库中 deepspeed/runtime/zero/stage_1_and_2.py 内以 reduce_scatter 为默认开关的完整实现,可以清晰看到一条一以贯之的设计主线:ZeRO 既分区存储,也分区通信。既然每个 rank 最终只需要属于自己的那份平均梯度,就用 reduce-scatter 精确投递、免去 all-reduce 中多余的 all-gather 阶段;再配合 reduce_bucket_size 约束桶内存、contiguous_gradients 对抗碎片化、overlap_comm 实现通信与计算重叠,最终让梯度平均的通信量与数据并行基线持平,并把通信耗时压低到官方公告所称的一半以内。对于希望深入 ZeRO 训练通信原理的读者,建议从 average_tensorreduce_ipg_grads 这两个函数入手研读源码,并结合 docs/_tutorials/zero.mddocs/_pages/config-json.md 的配置说明做小规模实验验证。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
898
5.82 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
921
1.84 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.8 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
531
596
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.02 K
519
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.36 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
391