示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载Join是 PyTorch 1.10 引入原型特性的通用上下文管理器专门用于解决分布式数据并行训练中各 rank 输入数量不均匀导致的进程挂起或报错问题。本文基于 PyTorch tutorials 仓库中的 advanced_source/generic_join.rst从DistributedDataParallelDDP与ZeroRedundancyOptimizerZeRO的联合使用出发完整讲解Join的用法、关键字参数传递方式、Joinable/JoinHook底层工作原理并带你实现一个与Join兼容的自定义玩具类。读完本文你将能够让训练循环直接兼容不均匀输入、在多个参与类之间共享上下文、以及为自己的类接入该机制。前置要求PyTorch 1.10Join从该版本起以原型特性引入API 后续可能变化。了解 DDP 的基础用法可参考仓库中的 intermediate_source/ddp_tutorial.rstGetting Started with Distributed Data Parallel。了解 ZeRO可参考仓库中的 recipes_source/zero_redundancy_optimizer.rstShard Optimizer States with ZeroRedundancyOptimizer。注意Join是 PyTorch 1.10 引入的原型prototype特性该 API 存在变动的可能。什么是Join不均匀输入问题的由来在 intermediate_source/ddp_tutorial.rst 的 Basic Use Case 一节中你可以看到使用 DDP 进行数据并行训练的标准骨架每个 rank 持有一份模型副本在每次 backward 过程中隐式地执行 all-reduce 来同步梯度。正如该教程 L160-L166 所总结的DDP 在 backward 阶段触发梯度同步并与 backward 计算重叠当backward()返回时param.grad已经包含同步后的梯度张量。这类集体通信collective communications要求进程组内所有 rank 都参与。因此一旦某个 rank 的输入更少提前耗尽数据其余 rank 在等待该 rank 参与通信时就会挂起或报错具体行为取决于后端如 NCCL/Gloo。更一般地任何在每次迭代中执行同步集体通信的类都会遇到同样的问题。Join正是为这一场景设计的上下文管理器它包裹在每 rank 的训练循环外层用于促进不均匀输入下的训练。其核心思想是先耗尽输入、提前加入join的 rank用影子通信shadow来顶替那些尚未加入的 rank 所执行的集体通信。具体如何顶替由各参与类提供的hooks钩子决定。从源码结构看Join的实现位于torch.distributed.algorithms.join模块本文所有示例均从torch.distributed.algorithms.join导入Join、Joinable、JoinHook。使用Join与DistributedDataParallelPyTorch 的DistributedDataParallel开箱即用地支持Join上下文管理器。下面是一个最小可运行的示例rank 1 比 rank 0 多一个输入若直接训练必然导致挂起而用Join包裹后两个 rank 都能正常跑完。import os import torch import torch.distributed as dist import torch.multiprocessing as mp from torch.distributed.algorithms.join import Join from torch.nn.parallel import DistributedDataParallel as DDP BACKEND nccl WORLD_SIZE 2 NUM_INPUTS 5 def worker(rank): os.environ[MASTER_ADDR] localhost os.environ[MASTER_PORT] 29500 dist.init_process_group(BACKEND, rankrank, world_sizeWORLD_SIZE) model DDP(torch.nn.Linear(1, 1).to(rank), device_ids[rank]) # Rank 1 gets one more input than rank 0 inputs [torch.tensor([1]).float() for _ in range(NUM_INPUTS rank)] num_inputs 0 with Join([model]): for input in inputs: num_inputs 1 loss model(input).sum() loss.backward() print(fRank {rank} has exhausted all {num_inputs} of its inputs!) def main(): mp.spawn(worker, nprocsWORLD_SIZE, joinTrue) if __name__ __main__: main()运行后输出如下rank 0 与 rank 1 的print()顺序可能任意交错Rank 0 has exhausted all 5 of its inputs! Rank 1 has exhausted all 6 of its inputs!关键点with Join([model]):接收一个参与类的列表。这里的model是 DDP 包装后的模块它在 backward 中执行 all-reduce因此必须被列在Join的参数中。历史背景在通用Join上下文管理器出现之前DistributedDataParallel自带一个join()上下文管理器。在上述示例中with Join([model]):与with model.join():是等价的。但旧的DistributedDataParallel.join()有一个明显局限它不允许同时存在多个参与类例如无法让DistributedDataParallel与ZeroRedundancyOptimizer一起参与。为什么 DDP 天然适配Join从 DDP 的工作原理可以理解Join的设计动机。intermediate_source/ddp_tutorial.rst 的 Skewed Processing Speeds 一节L171-L181明确指出DDP 的构造函数、forward 与 backward 都是分布式同步点。各进程应当发起相同次数的同步、以相同顺序到达同步点并大致同时进入。否则快的进程可能提前到达并在等待慢进程stragglers时超时。也就是说输入不均匀本质上属于处理速度倾斜的一种极端形态某个 rank 提前把输入处理完不再参与后续的同步。Join用影子通信替代提前加入者的真实通信从而让快的 rank 不必等待慢的 rank 也不会因无人响应而挂起。这也解释了为什么该教程建议在无法完全避免倾斜时应给init_process_group传入足够大的timeout值——而Join则从算法层面直接解决了提前耗尽这一最严重的不均匀情形。使用Join与DistributedDataParallelZeroRedundancyOptimizerJoin不仅能作用于单一类还能同时作用于多个类。PyTorch 的ZeroRedundancyOptimizer同样兼容该上下文管理器。下面在前例基础上加入 ZeROfrom torch.distributed.optim import ZeroRedundancyOptimizer as ZeRO from torch.optim import Adam def worker(rank): os.environ[MASTER_ADDR] localhost os.environ[MASTER_PORT] 29500 dist.init_process_group(BACKEND, rankrank, world_sizeWORLD_SIZE) model DDP(torch.nn.Linear(1, 1).to(rank), device_ids[rank]) optim ZeRO(model.parameters(), Adam, lr0.01) # Rank 1 gets one more input than rank 0 inputs [torch.tensor([1]).float() for _ in range(NUM_INPUTS rank)] num_inputs 0 # Pass both model and optim into Join() with Join([model, optim]): for input in inputs: num_inputs 1 loss model(input).sum() loss.backward() optim.step() print(fRank {rank} has exhausted all {num_inputs} of its inputs!)输出与前一节完全相同。唯一的实质变化是把ZeroRedundancyOptimizer实例也传入了Join()。这正是通用Join相对旧DDP.join()的核心优势——可以同时协调多个参与类。ZeRO 在Join中扮演的角色理解 ZeRO 的同步模式有助于把握它在Join中的行为。recipes_source/zero_redundancy_optimizer.rst 的 L40-L42 说明优化器的step()只更新其分片shard中的参数然后将更新后的参数广播给所有其他 DDP 进程使所有模型副本最终处于相同状态。这意味着 ZeRO 在step()中也要执行广播等集体通信属于每次迭代执行同步集体通信的类因此也必须纳入Join。同文档 L131-L137 还给出了 ZeRO 与普通 Adam 的峰值内存对比开启 ZeRO 时step()峰值内存约为普通 Adam 的一半说明 ZeRO 通过跨进程分片优化器状态来节省显存——而Join负责保证这种跨进程协作在不均匀输入下依然成立。向上下文管理器传递关键字参数参与类可以在运行时提供修改其在上下文管理器中行为的关键字参数。例如DistributedDataParallel提供了divide_by_initial_world_size参数用于决定梯度是除以初始 world size还是除以有效 world size即未加入的 rank 数量。这类参数可以直接传给上下文管理器with Join([model, optim], divide_by_initial_world_sizeFalse): for input in inputs: ...警告传给上下文管理器的关键字参数在所有参与类之间共享。这通常不构成限制——我们并不预期会出现多个Joinable需要对同一参数取不同值的情况但使用时仍需留意这一点。Join的工作原理Joinable、JoinHook与Join前面只是展示了用法接下来深入其内部实现。理解了Join类以及支撑它的Joinable、JoinHook三个类你就能完全掌握它的能力边界并为自己实现兼容类做好准备。Joinable兼容类的抽象基类要与Join上下文管理器兼容类必须继承抽象基类Joinable并实现以下方法join_hook(self, **kwargs) - JoinHook返回该Joinable对应的JoinHook实例。它决定了已加入joined的进程应如何影子化该Joinable在每轮迭代中执行的集体通信。join_device(self) - torch.device返回一个设备供Join上下文管理器执行集体通信使用例如torch.device(cuda:0)或torch.device(cpu)。join_process_group(self) - ProcessGroup返回供Join上下文管理器执行集体通信所用的进程组。其中join_device与join_process_group是必要属性确保上下文管理器能在已加入与未加入的进程之间调度集体通信。它们的典型用途有二每次迭代通过一个 all-reduce统计未加入进程的数量为throw_on_early_terminationTrue机制提供支撑下文详述。DistributedDataParallel与ZeroRedundancyOptimizer已经继承了Joinable并实现了上述方法所以在前面的示例中可以直接使用。此外Joinable子类应调用Joinable构造函数因为它会初始化一个JoinConfig实例供上下文管理器内部保证正确性该实例会以_join_config字段的形式保存在每个Joinable中。JoinHook两个入口钩子JoinHook为上下文管理器提供两个入口main_hook(self) - None只要还存在未加入的 rank每个已加入的 rank 就会反复调用该钩子。它的作用是在每个训练迭代例如一次 forward、一次 backward、一次优化器 step中影子化Joinable所执行的集体通信。post_hook(self, is_last_joiner: bool) - None当所有 rank 都加入后调用。它接收一个额外的bool参数is_last_joiner指示该 rank 是否为最后加入的 rank 之一。该参数可用于同步。以两个现成的实现为例ZeroRedundancyOptimizer的main hook会正常执行一次优化器 step因为已加入的 rank 仍要负责更新并同步自己持有的参数分片DistributedDataParallel的post hook会从某个最后加入的 rank 广播最终更新后的模型确保所有 rank 上的模型一致。Join把一切串起来最后看Join类本身。__init__(self, joinables: List[Joinable], enable: bool True, throw_on_early_termination: bool False)构造函数接收参与训练循环的Joinable列表——即那些每轮迭代都执行集体通信的类。其余两个参数enablebool默认True若已知输入不会不均匀可设为False此时上下文管理器退化为空操作行为类似contextlib.nullcontext()同时它也可能禁用参与Joinable中与 join 相关的计算。throw_on_early_terminationbool默认False设为True后一旦检测到不均匀输入每个 rank 都会立即抛异常。这适用于不满足上下文管理器要求的场景——最常见的是来自不同类的集体通信可能任意交错例如 DDP 搭配含SyncBatchNorm层的模型。此时应将该参数设为True由应用逻辑捕获异常并决定后续如何继续。核心逻辑位于__exit__()方法中只要还存在未加入的 rank就循环调用每个Joinable的main hook所有 rank 都加入后再调用它们的post hookmain hooks 与 post hooks 都按照传入Joinable的顺序依次执行。上下文管理器依赖来自未加入进程的心跳heartbeat。因此每个Joinable类应在自己的每轮迭代集体通信之前调用Join.notify_join_context()上下文管理器会确保只有传入的第一个Joinable真正发出心跳。警告如前面throw_on_early_termination所述Join上下文管理器与某些类组合不兼容。Joinable的JoinHook必须可串行执行——每个 hook 必须完整执行完才能进行下一个两个 hook 不能重叠。此外当前 main hooks 与 post hooks 都按同一确定性顺序迭代。如果这成为主要限制PyTorch 可能会修改 API 以允许自定义排序。实战让玩具类与Join兼容前面引入了多个概念下面用一个玩具示例把它们串起来。我们将实现一个类Counter用来统计在它所在 rank 加入之前所有 rank 总共处理了多少输入。这能直观展示如何让自定义类兼容Join。具体来说下面的代码让每个 rank 打印两行信息(1) 在它加入之前所有 rank 已处理的输入总数(2) 所有 rank 处理的输入总数。import os import torch import torch.distributed as dist import torch.multiprocessing as mp from torch.distributed.algorithms.join import Join, Joinable, JoinHook BACKEND nccl WORLD_SIZE 2 NUM_INPUTS 5 class CounterJoinHook(JoinHook): r Join hook for :class:Counter. Arguments: counter (Counter): the :class:Counter object using this hook. sync_max_count (bool): whether to sync the max count once all ranks join. def __init__( self, counter, sync_max_count ): self.counter counter self.sync_max_count sync_max_count def main_hook(self): r Shadows the counters all-reduce by all-reducing a dim-1 zero tensor. t torch.zeros(1, deviceself.counter.device) dist.all_reduce(t) def post_hook(self, is_last_joiner: bool): r Synchronizes the max count across all :class:Counter s if sync_max_countTrue. if not self.sync_max_count: return rank dist.get_rank(self.counter.process_group) common_rank self.counter.find_common_rank(rank, is_last_joiner) if rank common_rank: self.counter.max_count self.counter.count.detach().clone() dist.broadcast(self.counter.max_count, srccommon_rank) class Counter(Joinable): r Example :class:Joinable that counts the number of training iterations that it participates in. def __init__(self, device, process_group): super(Counter, self).__init__() self.device device self.process_group process_group self.count torch.tensor([0], devicedevice).float() self.max_count torch.tensor([0], devicedevice).float() def __call__(self): r Counts the number of inputs processed on this iteration by all ranks by all-reducing a dim-1 one tensor; increments its own internal count. Join.notify_join_context(self) t torch.ones(1, deviceself.device).float() dist.all_reduce(t) self.count t def join_hook(self, **kwargs) - JoinHook: r Return a join hook that shadows the all-reduce in :meth:__call__. This join hook supports the following keyword arguments: sync_max_count (bool, optional): whether to synchronize the maximum count across all ranks once all ranks join; default is False. sync_max_count kwargs.get(sync_max_count, False) return CounterJoinHook(self, sync_max_count) property def join_device(self) - torch.device: return self.device property def join_process_group(self): return self.process_group def find_common_rank(self, rank, to_consider): r Returns the max rank of the ones to consider over the process group. common_rank torch.tensor([rank if to_consider else -1], deviceself.device) dist.all_reduce(common_rank, opdist.ReduceOp.MAX, groupself.process_group) common_rank common_rank.item() return common_rank def worker(rank): assert torch.cuda.device_count() WORLD_SIZE os.environ[MASTER_ADDR] localhost os.environ[MASTER_PORT] 29500 dist.init_process_group(BACKEND, rankrank, world_sizeWORLD_SIZE) counter Counter(torch.device(fcuda:{rank}), dist.group.WORLD) inputs [torch.tensor([1]).float() for _ in range(NUM_INPUTS rank)] with Join([counter], sync_max_countTrue): for _ in inputs: counter() print(f{int(counter.count.item())} inputs processed before rank {rank} joined!) print(f{int(counter.max_count.item())} inputs processed across all ranks!) def main(): mp.spawn(worker, nprocsWORLD_SIZE, joinTrue) if __name__ __main__: main()由于 rank 0 处理 5 个输入、rank 1 处理 6 个输出为10 inputs processed before rank 0 joined! 11 inputs processed across all ranks! 11 inputs processed before rank 1 joined! 11 inputs processed across all ranks!解读rank 0 在耗尽自己的 5 个输入后提前加入此后 rank 1 每处理一个输入rank 0 都会通过 main hook 执行一次影子 all-reduce对 dim-1 的零张量做 all-reduce从而保持通信参与。最终 rank 0 记录到全部 11 个输入处理完成max_count而它在加入时已经看到的累计值是 105 个自己的 5 个 rank 1 的。这个示例的几个关键点正好对应前文的原理影子通信要形状匹配Counter实例每次迭代执行一个 all-reduce对 dim-1 的一张量因此 main hook 也执行一个 all-reduce对 dim-1 的零张量来顶替它保持通信次数与通信量一致。心跳位置正确Counter在__call__()开头调用Join.notify_join_context()因为这是其每轮迭代集体通信all-reduce之前的位置。若不调用上下文管理器将无法感知未加入进程仍在活跃。is_last_joiner用于确定广播源post hook 中find_common_rank(rank, is_last_joiner)通过一次ReduceOp.MAX的 all-reduce在所有 rank 中选出最后加入者中 rank 最大的那个作为广播源从而同步max_count。关键字参数透传我们把sync_max_countTrue传给上下文管理器它再转发给Counter的join_hook最终构造出对应的CounterJoinHook——这正是前文关键字参数在所有参与类之间共享的机制体现。小结与适用边界Join为 DDP、ZeRO 等每轮迭代执行同步集体通信的类提供了统一的不均匀输入解决方案用法极简把参与类列表传给with Join([...]):关键字参数直接透传支持多类协同DistributedDataParallel与ZeroRedundancyOptimizer可同时参与这是旧DDP.join()做不到的机制清晰提前加入的 rank 通过JoinHook.main_hook()影子化集体通信全部加入后由post_hook()完成最终同步如 DDP 的模型广播扩展容易继承Joinable、实现join_hook/join_device/join_process_group并在集体通信前调用Join.notify_join_context()即可让自定义类接入。需要记住的边界hooks 必须可串行执行、main/post hooks 按确定性顺序调用当不同类的集体通信会任意交错如 DDP SyncBatchNorm时应使用throw_on_early_terminationTrue交由应用层处理异常。此外Join目前仍是原型特性API 存在调整的可能升级 PyTorch 版本时建议留意相关变更说明。赞分享示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载相关推荐PyTorch 通用 Join 上下文管理器Generic Join Context Manager深入解析面向非均匀输入的分布式训练PyTorch 通用 Join 上下文管理器Generic Join Context Manager深入解析面向非均匀输入的分布式训练 导读 在 PyTo人工智能机器学习深度学习分布式训练模型编译终极指南如何使用ZeroRedundancyOptimizer优化PyTorch分布式训练内存消耗终极指南如何使用ZeroRedundancyOptimizer优化PyTorch分布式训练内存消耗 PyTorch作为深度学习领域的主流框架其分布式训练功能示例工程PyTorch DistributedDataParallel (DDP) 分布式训练原理深度解析PyTorch DistributedDataParallel DDP 分布式训练原理深度解析 导读 torch.nn.parallel.Distributed人工智能机器学习深度学习分布式训练模型编译上一篇IceCMS负载均衡配置高并发场景下的性能优化下一篇RIOT 中的 uTensor MNIST 手写数字识别示例从 TensorFlow 模型训练到 MCU 端推理创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
