进阶 pytorch.org 2026-10-07 23:07:35 · 6 阅读

第3章 使用 Join 上下文管理器处理不均衡输入的分布式训练

使用 Join 上下文管理器处理不均衡输入的分布式训练#创建时间:2021 年 8 月 4 日 | 最后更新:2025 年 9 月 3 日 | 最后验证:2024 年 11 月 5 日 作者:Andrew Gu

注意 你可以在 github 上查看和编辑本教程。

注意 Join 在 PyTorch 1.10 中作为原型功能引入,该 API 可能会发生变化。

在本教程中,你将了解:

Join 上下文管理器的概览。 如何在 DistributedDataParallel 中使用该上下文管理器的示例。 如何在 DistributedDataParallel 和 ZeroRedundancyOptimizer 中同时使用该上下文管理器的示例。 如何向上下文管理器传递关键字参数的示例。 深入剖析 Join 上下文管理器的工作原理。 一个演示如何让自定义类兼容该上下文管理器的示例。

前置要求#

PyTorch 1.10+ Getting Started with Distributed Data Parallel Shard Optimizer States with ZeroRedundancyOptimizer

什么是 Join?# 在 Getting Started with Distributed Data Parallel - Basic Use Case 中,你已经了解了使用 DistributedDataParallel 进行数据并行训练的基本框架。它会在每次反向传播中隐式调度 all-reduce,以同步各个 rank 之间的梯度。这类集合通信需要进程组中所有 rank 共同参与,如果某个 rank 的输入较少,其他 rank 就会挂起或报错(取决于后端)。更普遍地说,任何在每次迭代中执行同步集合通信的类都会遇到这个问题。 Join 是一个上下文管理器,用于包裹每个 rank 的训练循环,从而支持不均衡输入的训练。该上下文管理器让提前耗尽输入的 rank(即提前 join)能够“遮蔽”(shadow)尚未 join 的 rank 所执行的集合通信,具体的遮蔽方式由钩子(hook)指定。

在 DistributedDataParallel 中使用 Join# PyTorch 的 DistributedDataParallel 开箱即用地支持 Join 上下文管理器。下面是一个用法示例: 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, rank=rank, world_size=WORLD_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(f"Rank {rank} has exhausted all {num_inputs} of its inputs!")

def main(): mp.spawn(worker, nprocs=WORLD_SIZE, join=True)

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!

注意 在这个通用的 Join 上下文管理器推出之前,DistributedDataParallel 已经提供了自己的 join() 上下文管理器。在上面的例子中,with Join([model]): 等价于 with model.join():。现有 DistributedDataParallel.join() 的一个局限是不支持多个参与的类一起使用,比如 DistributedDataParallel 和 ZeroRedundancyOptimizer 的组合。

在 DistributedDataParallel 和 ZeroRedundancyOptimizer 中使用 Join# Join 上下文管理器不仅支持单个类,还支持多个类一起使用。PyTorch 的 ZeroRedundancyOptimizer 同样兼容该上下文管理器。下面我们来看如何修改上面的例子,让 DistributedDataParallel 和 ZeroRedundancyOptimizer 一起工作: from 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, rank=rank, world_size=WORLD_SIZE)

model = DDP(torch.nn.Linear(1, 1).to(rank), device_ids=[rank]) optim = ZeRO(model.parameters(), Adam, lr=0.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(f"Rank {rank} has exhausted all {num_inputs} of its inputs!")

输出与之前相同。关键变化就是把 ZeroRedundancyOptimizer 实例也传入了 Join()。

传递关键字参数# 参与类可以提供关键字参数,在运行时调整自己在上下文管理器中的行为。例如,DistributedDataParallel 提供了 divide_by_initial_world_size 参数,用于决定梯度是除以初始 world size 还是有效 world size(即尚未 join 的 rank 数量)。这类关键字参数可以直接传入上下文管理器。 with Join([model, optim], divide_by_initial_world_size=False): for input in inputs: ...

警告 传入上下文管理器的关键字参数对所有参与类是共享的。不过这应该不算限制,因为我们不预期会出现多个 Joinable 需要对同一参数使用不同设置的情况。尽管如此,这一点仍需留意。

Join 是如何工作的?# 看完这些初步示例后,我们深入探究 Join 上下文管理器的工作原理。这有助于你更全面地理解它的能力,也为让自定义类兼容它做好准备。这里我们会介绍 Join 类,以及配套的 Joinable 和 JoinHook 类。

Joinable# 首先,兼容 Join 上下文管理器的类必须继承抽象基类 Joinable。具体来说,Joinable 必须实现:

join_hook(self, kwargs) -> JoinHook

返回该 Joinable 的 JoinHook 实例,决定已 join 的进程如何遮蔽该 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 是必需的属性,以保证上下文管理器能在已 join 和未 join 的进程之间调度集合通信。一种用途是在每次迭代中用 all-reduce 统计未 join 的进程数量;另一种用途是实现 throw_on_early_termination=True 所需的机制(后文会解释)。 DistributedDataParallel 和 ZeroRedundancyOptimizer 已经继承了 Joinable 并实现了上述方法,这就是前面的例子中可以直接使用它们的原因。 Joinable 类应确保调用 Joinable 的构造函数,因为它会初始化一个 JoinConfig 实例,供上下文管理器内部使用以保证正确性。该实例会以 _join_config 字段的形式保存在每个 Joinable 中。

JoinHook# 接下来拆解 JoinHook 类。JoinHook 为上下文管理器提供两个入口:

main_hook(self) -> None

只要还存在未 join 的 rank,每个已 join 的 rank 就会反复调用该钩子。它用于遮蔽 Joinable 在每次训练迭代中执行的集合通信(如一次前向传播、反向传播和优化器 step)。

post_hook(self, is_last_joiner: bool) -> None

在所有 rank 都 join 之后调用一次。它会收到一个额外的布尔参数 is_last_joiner,表示该 rank 是否是最后 join 的之一,这个参数可用于同步。 举两个具体例子:ZeroRedundancyOptimizer 的 main 钩子照常执行一次优化器 step,因为已 join 的 rank 仍负责更新并同步自己那份参数分片;DistributedDataParallel 的 post 钩子则从最后 join 的某个 rank 广播最终更新后的模型,以保证所有 rank 上的模型一致。

Join# 最后看看这些组件如何整合进 Join 类本身。

__init__(self, joinables: List[Joinable], enable: bool = True, throw_on_early_termination: bool = False)

如前面的例子所示,构造函数接收一个参与训练循环的 Joinable 列表,即那些在每次迭代中执行集合通信的类。 enable 是一个布尔值,如果你确定输入不会不均衡,可以设为 False,此时上下文管理器就形同虚设,类似 contextlib.nullcontext()。这同时也会禁用参与 Joinable 中与 join 相关的计算。 throw_on_early_termination 是一个布尔值,设为 True 时,一旦检测到不均衡输入,每个 rank 都会立刻抛出异常。这适用于不符合上下文管理器要求的情况,最典型的就是来自不同类的集合通信可能任意交错时,例如 DistributedDataParallel 搭配包含 SyncBatchNorm 层的模型。此时应将该参数设为 True,让应用逻辑捕获异常并决定后续处理方式。

核心逻辑位于 __exit__() 方法中:只要还存在未 join 的 rank 就持续循环,调用每个 Joinable 的 main 钩子;待所有 rank 都 join 后,再调用它们的 post 钩子。main 钩子和 post 钩子都按 Joinable 传入的顺序依次执行。 上下文管理器需要来自未 join 进程的心跳信号。因此,每个 Joinable 类应在每次迭代的集合通信之前调用 Join.notify_join_context()。上下文管理器会确保只有第一个传入的 Joinable 实际发送心跳。

警告 如上文关于 throw_on_early_termination 的说明,Join 上下文管理器与某些类的组合不兼容。各 Joinable 的 JoinHook 必须是可串行执行的,因为每个钩子要完整执行完才会执行下一个,换句话说,两个钩子不能重叠。此外,目前 main 钩子和 post 钩子都按同样的确定性顺序遍历。如果这构成了重大限制,我们可能会修改 API 以支持自定义顺序。

让一个玩具类兼容 Join# 上一节介绍了不少概念,下面通过一个玩具示例来实践。我们要实现一个类,用于统计在本 rank join 之前所有 rank 看到的输入总数。这能让你对如何让自己的类兼容 Join 上下文管理器有一个基本认识。 具体来说,下面的代码让每个 rank 输出:(1) 在它 join 之前所有 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 counter's all-reduce by all-reducing a dim-1 zero tensor. """ t = torch.zeros(1, device=self.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_count=True``. """ 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, src=common_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], device=device).float() self.max_count = torch.tensor([0], device=device).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, device=self.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], device=self.device) dist.all_reduce(common_rank, op=dist.ReduceOp.MAX, group=self.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, rank=rank, world_size=WORLD_SIZE)

counter = Counter(torch.device(f"cuda:{rank}"), dist.group.WORLD) inputs = [torch.tensor([1]).float() for _ in range(NUM_INPUTS + rank)]

with Join([counter], sync_max_count=True): 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, nprocs=WORLD_SIZE, join=True)

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!

几个要点:

Counter 实例每次迭代只执行一次 all-reduce,因此 main 钩子也执行一次 all-reduce 来遮蔽它。 Counter 类在 __call__() 方法的开头调用了 Join.notify_join_context(),因为这是它每次迭代集合通信(即 all-reduce)之前的位置。 post 钩子中用 is_last_joiner 参数来确定广播源。 我们把 sync_max_count 关键字参数传给了上下文管理器,它随后会被转发给 Counter 的 join 钩子。

评论 (0)