Управляющий контекст объединения общего назначения
Управляющий контекст объединения общего назначения облегчает распределенное обучение на неравномерных входных данных. На этой странице описан API соответствующих классов: Join, Joinable, и JoinHook. Для получения учебного пособия см. Распределенное обучение с неравномерными входными данными с использованием управляющего контекста объединения.
-
class torch.distributed.algorithms.Join(joinables, enable=True, throw_on_early_termination=False, **kwargs)[source] -
Этот класс определяет управляющий контекст объединения общего назначения, который позволяет вызывать пользовательские хуки после того, как процесс присоединится. Эти хуки должны затенять коллективные коммуникации процессов, не присоединившихся, чтобы предотвратить зависание и ошибки, а также обеспечить корректность алгоритма. Подробности об определении хука см. в
JoinHook.Предупреждение
Управляющий контекст требует, чтобы каждый участвующий
Joinableвызывал методnotify_join_context()перед собственными коллективными коммуникациями на итерацию, чтобы обеспечить корректность.Предупреждение
Управляющий контекст требует, чтобы все
process_groupатрибуты в объектахJoinHookбыли одинаковыми. Если есть несколько объектовJoinHook, тоdeviceпервого используется. Для проверки отсутствия присоединившихся процессов и уведомления процессов об ошибке в случае включенияthrow_on_early_termination, используется информация о группе процессов и устройстве, обе операции выполняются с использованием all-reduce.- Параметры:
-
-
joinables (Список[Joinable]) – список участвующих
Joinable; их хуки перебираются в указанном порядке. -
enable (bool) – флаг, включающий обнаружение неравномерных входных данных; установка в
Falseотключает функциональность управляющего контекста и должна устанавливаться только в том случае, если пользователь знает, что входные данные не будут неравномерными (по умолчанию:True). -
throw_on_early_termination (bool) – флаг, управляющий тем, будет ли выбрасываться исключение при обнаружении неравномерных входных данных (по умолчанию:
False).
-
joinables (Список[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)[source] -
Уведомляет управляющий контекст объединения о том, что вызывающий процесс еще не присоединился; затем, если
throw_on_early_termination=True, проверяет, были ли обнаружены неравномерные входные данные (т. е. если один процесс уже присоединился), и выбрасывает исключение в случае их обнаружения.Этот метод должен вызываться из объекта
Joinableперед его коллективными коммуникациями на итерацию. Например, это должно вызываться в начале фазы forward вDistributedDataParallel.Только первый объект
Joinable, переданный в управляющий контекст, выполняет коллективные коммуникации в этом методе, а для других этот метод является пустым.- Параметры:
-
joinable (Joinable) – объект
Joinable, вызывающий этот метод. - Возвращаемое значение:
-
Дескриптор асинхронной операции all-reduce, предназначенный для уведомления управляющего контекста о том, что процесс еще не присоединился, если
joinableявляется первым переданным в управляющий контекст;Noneв противном случае.
-
class torch.distributed.algorithms.Joinable[source] -
Определяет абстрактный базовый класс для присоединяемых классов. Присоединяемый класс (наследуемый от
Joinable) должен реализовыватьjoin_hook(), который возвращает экземплярJoinHook, помимоjoin_device()иjoin_process_group(), которые возвращают информацию о устройстве и группе процессов соответственно.-
abstract property join_device: device -
Возвращает устройство, с которого следует выполнять коллективные коммуникации, необходимые для реализации управляющего контекста объединения.
-
abstract join_hook(**kwargs)[source]
-
abstract property join_process_group: Any -
Возвращает группу процессов для коллективных коммуникаций, необходимых для самой реализации управляющего контекста объединения.
-
-
class torch.distributed.algorithms.JoinHook[source] -
Это определяет крючок присоединения, который предоставляет две точки входа в менеджере контекста присоединения: основной крючок, который вызывается многократно, пока существует неподключённый процесс, и пост-крючок, который вызывается один раз, когда все процессы подключились.
Для реализации крючка присоединения для универсального менеджера контекста присоединения определите класс, который наследуется от
JoinHookи переопределитеmain_hook()иpost_hook()соответственно.-
main_hook()[source] -
Этот крючок вызывается многократно, пока существует неподключённый процесс, чтобы отслеживать коллективные коммуникации в одной итерации обучения (т. е. в одном проходе вперёд, обратном проходе и шаге оптимизатора).
-
post_hook(is_last_joiner)[source] -
Этот крючок вызывается после того, как все процессы подключились. Ему передаётся дополнительный
boolаргументis_last_joiner, который указывает, является ли ранг одним из последних, подключившихся.- Параметры:
-
is_last_joiner (bool) –
Trueесли ранг является одним из последних, подключившихся;Falseв противном случае.
-
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/distributed.algorithms.join.html