news 2026/9/9 19:58:25

PyTorch 分布式优化器详解:DistributedOptimizer、ZeroRedundancyOptimizer 与 PostLocalSGDOptimizer 使用指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch 分布式优化器详解:DistributedOptimizer、ZeroRedundancyOptimizer 与 PostLocalSGDOptimizer 使用指南

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暴露了三个对外成员:DistributedOptimizerPostLocalSGDOptimizerZeroRedundancyOptimizer,并通过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_mapregister_functional_optimas_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 的实现可以还原其工作流程:

  1. 按 worker 分组参数:构造时遍历params_rref,用param.owner()把参数按所在 worker 划分到per_worker_params_rref
  2. 选择本地优化器实现:若optimizer_class存在于functional_optim_map且 TorchScript 状态jit._state._enabled为真,则使用对应的功能式优化器构造器;否则回退到普通优化器并打印一条关于 GIL 的logger.warning
  3. 逐 worker 异步创建本地优化器:通过rpc.rpc_async(worker, optimizer_new_func, args=(optim_ctor, param_rrefs) + args, kwargs=kwargs)在参数所在 worker 上创建_LocalOptimizer(或_FunctionalLocalOptimizer)并拿到其RRef
  4. 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_classoptim.Optimizer在每个 worker 上实例化的优化器类
params_rreflist[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.distributedProcessGroup;默认取dist.group.WORLD(需先调用torch.distributed.init_process_group初始化)
parameters_as_bucket_view可选若为True,参数被打包进桶以加速通信,param.data指向桶视图的不同偏移;若为False,每个参数单独通信且param.data保持不变(默认False
overlap_with_ddp可选若为Truestep()与 DDP 的梯度同步重叠执行,此时parameters_as_bucket_view被忽略(默认False
**defaults可选其余关键字参数,透传给本地优化器(如lr=0.01
overlap_with_ddp 的启用前提

开启overlap_with_ddp=True需要同时满足:

  1. optimizer_class本身是功能式优化器,或存在功能式等价实现(即能在functional_optim_map中找到);
  2. 注册一个由ddp_zero_hook.py中的函数构造的 DDP 通信钩子;
  3. 参数被打包进与 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 实现层面的佐证

从源码可以看到若干内部机制的痕迹,帮助理解其行为:

  • 类同时继承OptimizerJoinable(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 参数与行为说明

参数类型含义
optimtorch.optim.Optimizer被包装的本地优化器
averageraveragers.ModelAverager执行 post-localSGD 算法的模型平均器实例

关键约定:PostLocalSGDOptimizer构造参数中warmup_steps的值必须与PostLocalSGDStatestart_localSGD_iter的值相同

5.3 模型平均器与 checkpoint 语义

  • 平均器可通过 torch/distributed/algorithms/model_averaging/averagers.py 中的PeriodicModelAverager(支持periodwarmup_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_groupsstatezero_grad(set_to_none=True)add_param_group()等接口,使用体验与普通torch.optim.Optimizer基本一致。

六、三大优化器的选型与适用场景小结

优化器依赖的分布式范式核心收益主要限制
DistributedOptimizerRPC + 分布式 autograd直接以RRef优化分散参数,天然支持参数服务器/模型并行;内部优先用功能式优化器规避 GIL当前不支持 CUDA 张量场景;跨客户端 step 与 forward 之间无全局顺序保证
ZeroRedundancyOptimizerDDP(collective)优化器状态分片约1/world_size,显著降低每 rank 峰值内存;支持与 DDP 重叠执行所有参数需同一种稠密类型;experimental;overlap 模式下前 2–3 个迭代不更新参数
PostLocalSGDOptimizerDDP + 通信钩子预热后以「子组梯度平均 + 周期全局模型平均」降低跨节点通信量warmup_stepsstart_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 场景实际权衡,建议先在小规模上对照开启/关闭ZeroRedundancyOptimizeroverlap_with_ddp的显存曲线,再决定是否引入后两种高级优化器。

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/9 19:58:20

MPU6050六轴传感器从入门到实战:驱动、校准与姿态解算全解析

简介:MPU6050.zip是一套面向ESP32开发者的MPU6050驱动与DMP姿态解算代码包,适合需要快速获取俯仰、翻滚、航偏角的嵌入式项目。资源共9个文件,以5个C头文件、3个C源文件和1个文本说明为主,包含MPU6050寄存器驱动、inv_mpu库及DMP运…

作者头像 李华
网站建设 2026/9/9 19:57:03

STM32串口奇偶校验实战:USART2配置与排坑指南

简介:面向STM32嵌入式开发者的一份实战工程,基于Cortex-M3内核的F103系列单片机,完整演示串口2带奇偶校验通信的配置方法。资源共78个文件,以头文件和C源码为主,包含串口、按键、LED、延时等模块驱动以及标准外设库&am…

作者头像 李华
网站建设 2026/9/9 19:55:49

如何避免“无标题”:打造高转化项目标题的实用指南

我注意到你提供的项目标题是“【无标题】”,这实际上是一个空的占位符,没有具体的项目名称或描述可供我展开分析。基于空标题无法生成有实质内容的博文。请重新提供有效的项目标题,并按照以下格式填写完整信息:项目标题: [一个有实…

作者头像 李华
网站建设 2026/9/9 19:55:04

Agno Workflow 后台执行实战:异步轮询与 WebSocket 实时事件流

Agno Workflow 后台执行实战:异步轮询与 WebSocket 实时事件流 【免费下载链接】agno Build, run, and manage agent platforms. 项目地址: https://gitcode.com/GitHub_Trending/ag/agno 导读 本篇技术指南以 cookbook/04_workflows/06_advanced_concepts/…

作者头像 李华