Spec-Zone.ru › PyTorch 2.14

Универсальный менеджер контекста Join

Создано: 6 июня 2025 г. | Последнее обновление: 6 июня 2025 г.

Универсальный менеджер контекста Join упрощает распределённое обучение на входных данных неодинакового размера. На этой странице описан API соответствующих классов: Join, Joinable и JoinHook. Учебное руководство см. в разделе Распределённое обучение на входных данных неодинакового размера с использованием менеджера контекста Join.

class torch.distributed.algorithms.Join(joinables, enable=True, throw_on_early_termination=False, **kwargs) [исходный код]

Этот класс определяет универсальный менеджер контекста Join, который позволяет вызывать пользовательские хуки после присоединения процесса.

Эти хуки должны имитировать коллективные коммуникации процессов, которые ещё не присоединились, чтобы избежать зависаний и ошибок, а также обеспечить корректность алгоритма. Подробнее об определении хука см. в разделе JoinHook.

Предупреждение

Для корректной работы менеджера контекста каждому участвующему объекту Joinable необходимо вызвать метод notify_join_context() перед собственными коллективными коммуникациями на каждой итерации.

Предупреждение

Менеджер контекста требует, чтобы все атрибуты process_group объектов JoinHook имели одинаковые значения. Если объектов JoinHook несколько, используется значение device первого объекта. Информация о группе процессов и устройстве используется для обнаружения процессов, которые ещё не присоединились, и для уведомления процессов о необходимости вызвать исключение, если включён параметр throw_on_early_termination. Для обоих действий используется all-reduce.

Параметры:
  • joinables (List[Joinable]) – список участвующих объектов Joinable; их хуки обрабатываются в указанном порядке.
  • enable (bool) – флаг, включающий обнаружение входных данных неодинакового размера; установка значения False отключает функциональность менеджера контекста и должна использоваться только в том случае, если пользователь уверен, что размеры входных данных будут одинаковыми (по умолчанию: True).
  • throw_on_early_termination (bool) – флаг, определяющий, следует ли вызывать исключение при обнаружении входных данных неодинакового размера (по умолчанию: False).

Пример:

>>> import os
>>> import torch
>>> import torch.distributed as dist
>>> import torch.multiprocessing as mp
>>> import torch.nn.parallel.DistributedDataParallel as DDP
>>> import torch.distributed.optim.ZeroRedundancyOptimizer as ZeRO
>>> from torch.distributed.algorithms.join import Join
>>>
>>> # On each spawned worker
>>> def worker(rank):
>>>     dist.init_process_group("nccl", rank=rank, world_size=2)
>>>     model = DDP(torch.nn.Linear(1, 1).to(rank), device_ids=[rank])
>>>     optim = ZeRO(model.parameters(), torch.optim.Adam, lr=0.01)
>>>     # Rank 1 gets one more input than rank 0
>>>     inputs = [torch.tensor([1.]).to(rank) for _ in range(10 + rank)]
>>>     with Join([model, optim]):
>>>         for input in inputs:
>>>             loss = model(input).sum()
>>>             loss.backward()
>>>             optim.step()
>>>     # All ranks reach here without hanging/erroring
static notify_join_context(joinable) [исходный код]

Уведомляет менеджер контекста Join о том, что вызывающий процесс ещё не присоединился.

Затем, если throw_on_early_termination=True, проверяет, обнаружены ли входные данные неодинакового размера (то есть присоединился ли уже один из процессов), и в таком случае вызывает исключение.

Этот метод следует вызывать из объекта Joinable перед его коллективными коммуникациями на каждой итерации. Например, его следует вызвать в начале прямого прохода в DistributedDataParallel.

Только первый объект Joinable, переданный менеджеру контекста, выполняет коллективные коммуникации в этом методе; для остальных вызов метода ничего не делает.

Параметры:

joinable (Joinable) – объект Joinable, вызывающий этот метод.

Возвращает:

Дескриптор асинхронной операции для all-reduce, предназначенной для уведомления менеджера контекста о том, что процесс ещё не присоединился, если joinable — первый объект, переданный менеджеру контекста; в противном случае — None.

class torch.distributed.algorithms.Joinable [исходный код]

Этот класс определяет абстрактный базовый класс для классов, поддерживающих Join.

Класс, поддерживающий Join (наследующий от Joinable), должен реализовать метод join_hook(), возвращающий экземпляр JoinHook, а также методы join_device() и join_process_group(), которые возвращают соответственно сведения об устройстве и группе процессов.

abstract property join_device: device

Возвращает устройство, с которого следует выполнять коллективные коммуникации, необходимые менеджеру контекста Join.

abstract join_hook(**kwargs) [исходный код]

Возвращает экземпляр JoinHook для указанного объекта Joinable.

Параметры:

kwargs (dict) – словарь dict с именованными аргументами, изменяющими поведение хука Join во время выполнения; всем экземплярам Joinable, использующим один и тот же менеджер контекста Join, передаётся одно и то же значение kwargs.

Тип возвращаемого значения:

JoinHook

abstract property join_process_group: Any

Возвращает группу процессов для коллективных коммуникаций, необходимых самому менеджеру контекста Join.

class torch.distributed.algorithms.JoinHook [исходный код]

Этот класс определяет хук Join, предоставляющий две точки входа в менеджере контекста Join.

Точки входа: основной хук, многократно вызываемый, пока остаются процессы, которые ещё не присоединились, и постхук, вызываемый после присоединения всех процессов.

Чтобы реализовать хук Join для универсального менеджера контекста Join, определите класс, наследующий от JoinHook, и при необходимости переопределите main_hook() и post_hook().

main_hook() [исходный код]

Вызывайте этот хук, пока остаются процессы, которые ещё не присоединились, чтобы имитировать коллективные коммуникации на итерации обучения.

Итерация обучения — это один прямой проход, обратный проход и шаг оптимизатора.

post_hook(is_last_joiner) [исходный код]

Вызывайте хук после присоединения всех процессов.

Ему передаётся дополнительный аргумент bool is_last_joiner, указывающий, является ли ранг одним из последних присоединившихся.

Параметры:

is_last_joiner (bool) – True, если ранг является одним из последних присоединившихся; в противном случае — False.

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/distributed.algorithms.join.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API