PyTorch 分布式优化器详解:DistributedOptimizer、ZeroRedundancyOptimizer 与 PostLocalSGDOptimizer 使用指南
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
导读
torch.distributed.optim是 PyTorch 面向分布式训练(RPC 远程参数、DDP 数据并行、模型并行等场景)提供的分布式优化器模块。本文围绕该模块的官方 API 文档展开,结合仓库源码(torch/distributed/optim)逐层拆解三个核心类的原理与用法,帮助读者掌握:通过 DistributedOptimizer 直接优化分布在多个 worker 上的远程参数、通过 ZeroRedundancyOptimizer 分片优化器状态以降低显存峰值、通过 PostLocalSGDOptimizer 实现通信后本地 SGD,并学会用register_functional_optim接入自定义的功能式优化器。
::: {warning} 官方文档特别提示:Distributed Optimizer(即DistributedOptimizer)当前不支持 CUDA 张量场景(Distributed optimizer is not currently supported when using CUDA tensors),请在 CPU 分布式场景或明确支持的运行方式下使用。 :::
一、模块全景:torch.distributed.optim 提供什么
官方 API 文档(docs/source/distributed.optim.md)通过automodule:: torch.distributed.optim暴露了三个对外成员:DistributedOptimizer、PostLocalSGDOptimizer、ZeroRedundancyOptimizer,并通过autofunction:: register_functional_optim暴露了模块工具层的注册接口。
从 torch/distributed/optim/init.py 可以看到模块内部结构与导出清单(__all__):
DistributedOptimizer:面向RPC 分布式 autograd的优化器,接收远程参数引用RRef,在参数所在的 worker 上本地执行优化;PostLocalSGDOptimizer:包装任意本地优化器,在预热(warm-up)后周期性执行全局模型平均,配合 DDP 通信钩子实现 post-local SGD 算法;ZeroRedundancyOptimizer:按 ZeRO 思想将优化器状态分片到各 rank,每个 rank 只维护约1/world_size的优化器状态;- 一组内部
_Functional*优化器(_FunctionalAdadelta、_FunctionalAdam、_FunctionalAdamW、_FunctionalSGD等)与functional_optim_map、register_functional_optim、as_functional_optim工具函数,见 torch/distributed/optim/utils.py。
值得注意的实现细节:DistributedOptimizer的导入被if hasattr(torch._C, "_rpc_init")守卫(见 torch/distributed/optim/init.py),即其可用性依赖于torch.distributed.rpc是否编译进当前构建。
二、DistributedOptimizer:直接优化远程参数的分布式优化器
DistributedOptimizer面向的是 RPC 风格的分布式训练(如分布式模型并行 / 参数服务器)。它接收一批散布在不同 worker 上的参数引用,并让每个 worker 用本地优化器就地更新属于自己的那部分参数。
2.1 设计原理
从 torch/distributed/optim/optimizer.py 的实现可以还原其工作流程:
- 按 worker 分组参数:构造时遍历
params_rref,用param.owner()把参数按所在 worker 划分到per_worker_params_rref; - 选择本地优化器实现:若
optimizer_class存在于functional_optim_map且 TorchScript 状态jit._state._enabled为真,则使用对应的功能式优化器构造器;否则回退到普通优化器并打印一条关于 GIL 的logger.warning; - 逐 worker 异步创建本地优化器:通过
rpc.rpc_async(worker, optimizer_new_func, args=(optim_ctor, param_rrefs) + args, kwargs=kwargs)在参数所在 worker 上创建_LocalOptimizer(或_FunctionalLocalOptimizer)并拿到其RRef; - step 阶段:
step(context_id)先通过dist_autograd._is_valid_context(context_id)校验 autograd 上下文,再对每个远程优化器发起rpc.rpc_async(optimizer.owner(), optimizer_step_func, args=(optimizer, context_id)),最后_wait_for_all(rpc_futs)阻塞直到所有 worker 完成更新。
本地优化器通过dist_autograd.get_gradients(autograd_ctx_id)拉取该上下文中与本地参数相关的梯度。普通路径(_LocalOptimizer)会把梯度写入param.grad后调用self.optim.step(),并用一把进程级global_lock串行化同一 worker 上并发到来的 step(见 optimizer.py);功能式路径(_FunctionalLocalOptimizer)则直接以「梯度列表」调用功能式优化器的step(grads),全程不写param.grad、不持有全局锁(见 optimizer.py)。
2.2 完整示例
官方文档给出了如下端到端示例(见 optimizer.py),演示「forward → backward → optimizer step」三段式用法:
import torch.distributed.autograd as dist_autograd import torch.distributed.rpc as rpc from torch import optim from torch.distributed.optim import DistributedOptimizer with dist_autograd.context() as context_id: # Forward pass:在 worker1 上执行远程计算,返回参数 RRef rref1 = rpc.remote("worker1", torch.add, args=(torch.ones(2), 3)) rref2 = rpc.remote("worker1", torch.add, args=(torch.ones(2), 1)) loss = rref1.to_here() + rref2.to_here() # Backward pass:基于 context 做分布式反向传播 dist_autograd.backward(context_id, [loss.sum()]) # Optimizer:针对远程参数 RRef 创建 SGD 优化器 dist_optim = DistributedOptimizer( optim.SGD, [rref1, rref2], lr=0.05, ) dist_optim.step(context_id)2.3 构造参数与行为约定
DistributedOptimizer.__init__(optimizer_class, params_rref, *args, **kwargs)的参数约定如下:
| 参数 | 类型 | 含义 |
|---|---|---|
optimizer_class | optim.Optimizer | 在每个 worker 上实例化的优化器类 |
params_rref | list[RRef] | 待优化参数的远程引用列表,可指向本地或远程参数 |
args/kwargs | 可变参数 | 透传给每个 worker 上优化器构造函数的参数 |
其step(context_id)行为需要特别注意(见类 docstring):
- 同一或不同客户端对
step的并发调用会在每个 worker 上被串行化——因为每个 worker 的优化器同一时刻只能处理一份梯度; - 但不保证某个客户端的「forward-backward-optimizer」完整序列与其他客户端互斥执行,因此被应用的梯度未必对应某个 worker 上最近一次 forward 的结果;
- 跨 worker 之间没有保证的执行顺序。
2.4 功能式优化器带来的并发收益
DistributedOptimizer内部在可用时会优先使用功能式优化器(源码注释:"DistributedOptimizer uses a functional optimizer internally when one is available...so that optimizer updates are not blocked by the Python Global Interpreter Lock (GIL)")。这对多线程训练场景(例如分布式模型并行)意义重大:参数更新不再被 GIL 阻塞。目前这一特性已覆盖大多数常用优化器。
functional_optim_map(utils.py)内置的映射为:
| 用户传入的优化器类 | 内部使用的功能式类 |
|---|---|
optim.SGD | _FunctionalSGD |
optim.Adam | _FunctionalAdam |
optim.AdamW | _FunctionalAdamW |
optim.Adagrad | _FunctionalAdagrad |
optim.Adadelta | _FunctionalAdadelta |
optim.RMSprop | _FunctionalRMSprop |
optim.Rprop | _FunctionalRprop |
optim.Adamax | _FunctionalAdamax |
这些功能式实现对应文件即 functional_sgd.py、functional_adam.py、functional_adamw.py 等。若某优化器类不在映射内,DistributedOptimizer会回退到普通_LocalOptimizer,并给出警告:在多线程环境(如 CPU 上的分布式模型并行)下可能因 GIL 导致计算变慢。
三、register_functional_optim:注册自定义功能式优化器
torch.distributed.optim.utils.register_functional_optim(key, optim)是模块对外暴露的唯一工具函数接口,用于向functional_optim_map插入新的功能式优化器。
3.1 签名与约束
从 utils.py 源码看,其行为是:仅当key尚不在functional_optim_map中时才写入(if key not in functional_optim_map)。官方文档特别强调:key 与 optimizer 不需要是torch.optim.Optimizer类型(例如可以是自定义优化器),因此该接口对第三方/自定义优化器同样开放。
3.2 使用示例
# import the new functional optimizer(以模块方式引入) from xyz import fn_optimizer from torch.distributed.optim.utils import register_functional_optim fn_optim_key = "XYZ_optim" register_functional_optim(fn_optim_key, fn_optimizer)3.3 配套工具 as_functional_optim
除注册接口外,utils.py 还提供了as_functional_optim(optim_cls, *args, **kwargs):它把用户传入的优化器类在functional_optim_map中查找其功能式对应物,找不到时抛出ValueError(f"Optimizer {optim_cls} does not have a functional counterpart!");找到后以空参数列表 +_allow_empty_param_list=True实例化功能式优化器。参数列表之所以可为空,是因为功能式优化器的参数在每次 step 时才通过参数列表显式传入。
四、ZeroRedundancyOptimizer:优化器状态的 ZeRO 分片
ZeroRedundancyOptimizer是本文档中应用面最广的成员,用于降低 DDP 训练中每个 rank 的峰值显存/内存占用。
4.1 核心思想
从 zero_redundancy_optimizer.py 的 docstring 可以提炼它的运行机理:
- 包装任意
torch.optim.Optimizer,将优化器状态按 ZeRO 论文(arXiv:1910.02054)的方式跨进程组分片; - 每个 rank 上的本地优化器只负责更新约
1/world_size的参数,因而只需要维护1/world_size的优化器状态(如 Adam 的 m、v 动量缓冲); - 本 rank 完成本地更新后,将本分片参数广播给组内所有对等进程,使各模型副本保持一致;
- 与
torch.nn.parallel.DistributedDataParallel配合使用可显著降低每 rank 峰值内存。
4.2 参数切分策略
官方文档明确了两条切分事实:
- ZeroRedundancyOptimizer 采用sorted-greedy(排序贪心)算法为每个 rank 打包一批参数;
- 每个参数完整地归属某一个 rank,不会把一个参数再跨 rank 切分;
- 该划分是任意的,可能与参数注册或使用顺序不一致。
4.3 构造参数详解
构造函数签名与参数(见 zero_redundancy_optimizer.py):
ZeroRedundancyOptimizer( params, optimizer_class, process_group=None, parameters_as_bucket_view=False, overlap_with_ddp=False, **defaults, )| 参数 | 必选/可选 | 含义与默认值 |
|---|---|---|
params | 必选 | Iterable[torch.Tensor]或Iterable[dict],给出将被跨 rank 分片的所有参数 |
optimizer_class | 必选 | 本地优化器类(torch.nn.Optimizer) |
process_group | 可选 | torch.distributed的ProcessGroup;默认取dist.group.WORLD(需先调用torch.distributed.init_process_group初始化) |
parameters_as_bucket_view | 可选 | 若为True,参数被打包进桶以加速通信,param.data指向桶视图的不同偏移;若为False,每个参数单独通信且param.data保持不变(默认False) |
overlap_with_ddp | 可选 | 若为True,step()与 DDP 的梯度同步重叠执行,此时parameters_as_bucket_view被忽略(默认False) |
**defaults | 可选 | 其余关键字参数,透传给本地优化器(如lr=0.01) |
overlap_with_ddp 的启用前提
开启overlap_with_ddp=True需要同时满足:
optimizer_class本身是功能式优化器,或存在功能式等价实现(即能在functional_optim_map中找到);- 注册一个由
ddp_zero_hook.py中的函数构造的 DDP 通信钩子; - 参数被打包进与 DDP 一致的桶中——这正是
parameters_as_bucket_view被忽略的原因。
源码中(zero_redundancy_optimizer.py),overlap_with_ddp=True时本地优化器的初始化被延迟到运行时,待 DDP 完成梯度桶重建、收集到分桶信息(_OverlapInfo)后再初始化;同时若parameters_as_bucket_view=True,会发出警告提示该参数将被忽略。
4.4 完整示例
import torch import torch.nn as nn from torch.distributed.optim import ZeroRedundancyOptimizer from torch.nn.parallel import DistributedDataParallel as DDP model = nn.Sequential(*[nn.Linear(2000, 2000).to(rank) for _ in range(20)]) ddp = DDP(model, device_ids=[rank]) opt = ZeroRedundancyOptimizer( ddp.parameters(), optimizer_class=torch.optim.Adam, lr=0.01 ) ddp(inputs).sum().backward() opt.step()注意:示例中传入的是ddp.parameters()——参数先经 DDP 包装,再由 ZeRO 优化器接管分片与广播,二者各司其职(DDP 负责梯度同步,ZeRO 负责优化器状态分片)。
4.5 使用限制与警告(务必阅读)
- 类型限制:当前要求传入的所有参数是相同的稠密类型("all of the passed-in parameters are the same dense type");
- 前几个迭代不更新参数:开启
overlap_with_ddp=True时,由于需要等待 DDP 分桶信息定型——static_graph=False时直到第二次 forward、static_graph=True时直到第三次 forward——训练的前两到三个迭代不会真正执行参数更新。规避手段之一是在前面补几个 dummy 输入; - 实验性:ZeroRedundancyOptimizer 处于 experimental 状态,接口后续可能变化。
4.6 实现层面的佐证
从源码可以看到若干内部机制的痕迹,帮助理解其行为:
- 类同时继承
Optimizer与Joinable(zero_redundancy_optimizer.py),配合_ZeROJoinHook支持不均匀输入下的 Join 语义; _OverlapStatus枚举(UNINITIALIZED → DDP_HAS_REBUILT_BUCKETS → INITIALIZED)刻画了与 DDP 重叠时优化器的三阶段初始化状态机(见 zero_redundancy_optimizer.py);- 每轮迭代结束时需
wait_for_broadcasts()等待参数广播完成、clear_per_iter_info()清理按迭代变化的缓存结构(见 zero_redundancy_optimizer.py)。
在仓库测试中可找到大量用法参照,例如 test/distributed/optim/test_zero_redundancy_optimizer.py 中的各类组合测试(含 DDP 组合、Join 场景、checkpoint 恢复等)。
五、PostLocalSGDOptimizer:通信后本地 SGD
PostLocalSGDOptimizer包装任意torch.optim.Optimizer,实现post-local SGD算法(论文 arXiv:1808.07217)。它的工作方式为:
- 每一步都运行本地优化器 step;
- 预热阶段结束后,周期性地在本地优化器应用之后,对所有参数执行一次全局模型平均。
5.1 使用步骤与参数
官方示例给出了完整的接线方式(见 post_localSGD_optimizer.py),核心步骤是:构造 DDP → 注册 post-localSGD 通信钩子 → 用PostLocalSGDState+post_localSGD_hook→ 创建包装优化器 → 正常训练循环。
import torch import torch.distributed as dist import torch.distributed.algorithms.model_averaging.averagers as averagers import torch.nn as nn from torch.distributed.optim import PostLocalSGDOptimizer from torch.distributed.algorithms.ddp_comm_hooks.post_localSGD_hook import ( PostLocalSGDState, post_localSGD_hook, ) model = nn.parallel.DistributedDataParallel( module, device_ids=[rank], output_device=rank ) # 注册 post-localSGD 通信钩子 state = PostLocalSGDState(process_group=None, subgroup=None, start_localSGD_iter=100) model.register_comm_hook(state, post_localSGD_hook) # 创建 post-localSGD 优化器,包装本地优化器。 # 注意:PostLocalSGDOptimizer 的 warmup_steps 必须与 # PostLocalSGDState 的 start_localSGD_iter 保持一致。 local_optim = torch.optim.SGD(params=model.parameters(), lr=0.01) opt = PostLocalSGDOptimizer( optim=local_optim, averager=averagers.PeriodicModelAverager(period=4, warmup_steps=100) ) # 前 100 步:DDP 每一步做全局梯度平均。 # 100 步之后:DDP 在每个子组(默认节点内)内做梯度平均, # 而 post-localSGD 优化器在应用本地优化器后,每 4 步做一次全局模型平均。 for step in range(0, 200): opt.zero_grad() loss = loss_fn(output, labels) loss.backward() opt.step()训练过程中会自然发生两阶段切换:前start_localSGD_iter(这里为 100)步 DDP 做全局梯度 all-reduce;之后退化为子组内梯度平均 + 周期性全局模型平均(此处每 4 步一次),从而在保证收敛质量的同时大幅削减跨节点的通信量。
5.2 参数与行为说明
| 参数 | 类型 | 含义 |
|---|---|---|
optim | torch.optim.Optimizer | 被包装的本地优化器 |
averager | averagers.ModelAverager | 执行 post-localSGD 算法的模型平均器实例 |
关键约定:PostLocalSGDOptimizer构造参数中warmup_steps的值必须与PostLocalSGDState中start_localSGD_iter的值相同。
5.3 模型平均器与 checkpoint 语义
- 平均器可通过 torch/distributed/algorithms/model_averaging/averagers.py 中的
PeriodicModelAverager(支持period、warmup_steps)等实现提供; - 从源码(post_localSGD_optimizer.py)可见,其
step()的语义是先执行self.optim.step(),随后调用self.averager.average_parameters(params=self.param_groups); - checkpoint 语义:
state_dict()在底层优化器状态之外,额外写入一个"step"键记录平均器当前步数,确保重载后不会重复触发无谓的预热;load_state_dict()会恢复平均器步数,若 state dict 中没有"step"项则发出警告并把平均器步数置 0; - 该类还透传/代理了
param_groups、state、zero_grad(set_to_none=True)、add_param_group()等接口,使用体验与普通torch.optim.Optimizer基本一致。
六、三大优化器的选型与适用场景小结
| 优化器 | 依赖的分布式范式 | 核心收益 | 主要限制 |
|---|---|---|---|
DistributedOptimizer | RPC + 分布式 autograd | 直接以RRef优化分散参数,天然支持参数服务器/模型并行;内部优先用功能式优化器规避 GIL | 当前不支持 CUDA 张量场景;跨客户端 step 与 forward 之间无全局顺序保证 |
ZeroRedundancyOptimizer | DDP(collective) | 优化器状态分片约1/world_size,显著降低每 rank 峰值内存;支持与 DDP 重叠执行 | 所有参数需同一种稠密类型;experimental;overlap 模式下前 2–3 个迭代不更新参数 |
PostLocalSGDOptimizer | DDP + 通信钩子 | 预热后以「子组梯度平均 + 周期全局模型平均」降低跨节点通信量 | warmup_steps与start_localSGD_iter必须一致;需要配套的 model averager |
三者同属 torch/distributed/optim 模块、共享functional_optim_map这一功能式优化器基础设施,但服务于不同的分布式训练拓扑:DistributedOptimizer面向 RPC 远程参数,ZeroRedundancyOptimizer面向 DDP 内存优化,PostLocalSGDOptimizer面向 DDP 通信量优化。
七、进一步阅读
- 官方 API 文档出处:docs/source/distributed.optim.md
- 模块导出与结构:torch/distributed/optim/init.py
- 远程参数优化器实现:torch/distributed/optim/optimizer.py
- ZeRO 分片优化器实现:torch/distributed/optim/zero_redundancy_optimizer.py
- post-local SGD 优化器实现:torch/distributed/optim/post_localSGD_optimizer.py
- 功能式优化器注册与映射:torch/distributed/optim/utils.py
- 功能式实现集合:functional_sgd.py、functional_adam.py、functional_adamw.py 等
- 相关测试参考:test/distributed/optim/test_zero_redundancy_optimizer.py、test/distributed/optim/test_apply_optimizer_in_backward.py、test/distributed/optim/test_named_optimizer.py
如需在 DDP 场景实际权衡,建议先在小规模上对照开启/关闭ZeroRedundancyOptimizer与overlap_with_ddp的显存曲线,再决定是否引入后两种高级优化器。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考