Управляющий контекстом объединения общего назначения
Управляющий контекстом объединения общего назначения облегчает распределенное обучение на неравномерных входных данных. На этой странице описан 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перед его коллективными коммуникациями на итерацию. Например, это следует вызывать в начале прохода вперёд в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/2.1/distributed.algorithms.join.html