DeepSpeed ZeRO 通信优化深度解析:用 reduce-scatter 分区感知规约替代 all-reduce
本文基于仓库内发布于 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_scatter(stage_1_and_2.py#L230)。同一文件中用 partition_gradients 区分两个阶段:True 时为 ZeRO-2(梯度也被分区),False 时为 ZeRO-1(仅优化器状态分区),见 stage_1_and_2.py#L224-L226。可见"分区感知、只保留本地分区所需梯度切片"的逻辑对两个阶段是统一的。
2. average_tensor:一个方法内的两条规约路径
梯度规约的核心函数是 average_tensor(stage_1_and_2.py#L1360-L1478):
- 关闭 reduce_scatter 时,走
gradient_reduction_w_predivide(stage_1_and_2.py#L1269-L1298),内部对整桶梯度调用dist.all_reduce,并通过gradient_predivide_factor、postscale_gradients、gradient_average等参数在 fp16 下做"先除后归约"以控制数值稳定性; - 开启 reduce_scatter 时,代码遍历桶内每个参数的分区元数据
param_to_partition_ids与grad_start_offset,把梯度张量按目标分区(dst_rank, bucket_offset, numel)切成若干连续切片(stage_1_and_2.py#L1386-L1442),随后按切片目标分组:- 切片只属于单一目标 rank 时,使用
("reduce", dst, process_group)为键、对该目标做dist.reduce(经allreduce_no_retain→allreduce_bucket(rank=dst),其内部即 stage_1_and_2.py#L1866-L1871 的dist.reduce到目标全局 rank); - 存在多 rank 副本需求等特殊场景时,才回退到
allreduce_and_scatter路径(stage_1_and_2.py#L1466-L1470)。
- 切片只属于单一目标 rank 时,使用
换言之,一个参数可能横跨多个分区、一个"规约桶"里也往往有多个连续片段指向同一目标 rank,代码会将这些片段合并后再一次性规约(stage_1_and_2.py#L1435-L1442),尽量让每次集合通信都有足够大的数据量。
3. 归约结果直接对位本地分区,非本地梯度尽早释放
在 reduce_ipg_grads(stage_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-L536 与 deepspeed/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_factor与postscale_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_tensor 与 reduce_ipg_grads 这两个函数入手研读源码,并结合 docs/_tutorials/zero.md 与 docs/_pages/config-json.md 的配置说明做小规模实验验证。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0629
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
HivisionIDPhotos⚡️HivisionIDPhotos: a lightweight and efficient AI ID photos tools. 一个轻量级的AI证件照制作算法。Python07
DragonOSDragonOS is an operating system developed from scratch using Rust, with Linux compatibility. It is designed for **Serverless** scenarios. 使用Rust从0自研内核,具有Linux兼容性的操作系统,面向云计算Serverless场景而设计。Rust00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00