Универсальный менеджер контекста 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).
-
joinables (List[Joinable]) – список участвующих объектов
Пример:
>>> 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.
-
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)[исходный код] -
Вызывайте хук после присоединения всех процессов.
Ему передаётся дополнительный аргумент
boolis_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