DistributedDataParallel
-
class torch.nn.parallel.DistributedDataParallel(module, device_ids=None, output_device=None, dim=0, broadcast_buffers=None, init_sync=True, process_group=None, bucket_cap_mb=None, find_unused_parameters=False, check_reduction=False, gradient_as_bucket_view=False, static_graph=False, delay_all_reduce_named_params=None, param_to_hook_all_reduce=None, mixed_precision=None, device_mesh=None, skip_all_reduce_unused_params=False, bucket_cap_mb_list=None, batched_grad_copy=False, forward_sync_buffers=None)[исходный код] -
Реализует распределённый параллелизм данных на уровне модуля на основе
torch.distributed.Этот контейнер обеспечивает параллелизм данных, синхронизируя градиенты между репликами модели. Устройства, между которыми выполняется синхронизация, задаются входным параметром
process_group, который по умолчанию включает весь мир. Обратите внимание, чтоDistributedDataParallelне разбивает входные данные на части и не распределяет их иным образом между участвующими GPU; пользователь должен самостоятельно определить, как это сделать, например, с помощьюDistributedSampler.См. также: Основы и Используйте nn.parallel.DistributedDataParallel вместо multiprocessing или nn.DataParallel. Применяются те же ограничения на входные данные, что и в
torch.nn.DataParallel.Для создания этого класса необходимо, чтобы
torch.distributedуже была инициализирована вызовомtorch.distributed.init_process_group().Доказано, что
DistributedDataParallelзначительно быстрее, чемtorch.nn.DataParallel, при обучении с параллелизмом данных на нескольких GPU одного узла.Чтобы использовать
DistributedDataParallelна хосте с N GPU, необходимо запуститьNпроцессов, гарантируя, что каждый процесс работает исключительно с одним GPU с индексом от 0 до N-1. Это можно сделать, задавCUDA_VISIBLE_DEVICESдля каждого процесса или вызвав следующий API для GPU:>>> torch.cuda.set_device(i)
или вызвав унифицированный API для ускорителя:
>>> torch.accelerator.set_device_index(i)
где i принимает значения от 0 до N-1. В каждом процессе для создания этого модуля следует использовать следующий код:
>>> if torch.accelerator.is_available(): >>> device_type = torch.accelerator.current_accelerator().type >>> vendor_backend = torch.distributed.get_default_backend_for_device(device_type) >>> >>> torch.distributed.init_process_group( >>> backend=vendor_backend, world_size=N, init_method='...' >>> ) >>> model = DistributedDataParallel(model, device_ids=[i], output_device=i)
Также можно использовать новейший API для инициализации:
>>> torch.distributed.init_process_group(device_id=i)
Чтобы запустить несколько процессов на одном узле, можно использовать
torch.distributed.launchилиtorch.multiprocessing.spawn.Примечание
Краткое введение во все функции, связанные с распределённым обучением, см. в документе Обзор распределённых возможностей PyTorch.
Примечание
DistributedDataParallelможно использовать вместе сtorch.distributed.optim.ZeroRedundancyOptimizer, чтобы уменьшить объём памяти, занимаемой состояниями оптимизатора на каждом ранге. Подробности см. в рецепте ZeroRedundancyOptimizer.Примечание
В настоящее время backend
nccl— самый быстрый и настоятельно рекомендуемый backend при использовании GPU. Это относится к распределённому обучению как на одном узле, так и на нескольких узлах.Примечание
Этот модуль также поддерживает распределённое обучение со смешанной точностью. Это означает, что модель может иметь параметры разных типов, например смешанные типы
fp16иfp32; сведение градиентов для таких параметров работает корректно.Примечание
Если в одном процессе для создания контрольной точки модуля используется
torch.save, а в других процессах для восстановления —torch.load, убедитесь, чтоmap_locationправильно настроен для каждого процесса. Безmap_locationtorch.loadвосстановит модуль на устройствах, с которых он был сохранён.Примечание
Если модель обучается на
Mузлах с использованиемbatch=N, градиент будет вMраз меньше по сравнению с той же моделью, обучаемой на одном узле с использованиемbatch=M*N, если функция потерь суммируется (а НЕ усредняется, как обычно) по примерам в пакете (поскольку градиенты между разными узлами усредняются). Учитывайте это, если хотите получить математически эквивалентный процесс обучения по сравнению с локальным обучением. Однако в большинстве случаев модель, обёрнутую в DistributedDataParallel, модель, обёрнутую в DataParallel, и обычную модель на одном GPU можно считать эквивалентными (например, использовать одинаковую скорость обучения для эквивалентного размера пакета).Примечание
Параметры никогда не рассылаются между процессами. Модуль выполняет операцию all-reduce над градиентами и предполагает, что во всех процессах оптимизатор изменяет их одинаковым образом. Буферы (например, статистики BatchNorm) в каждой итерации передаются из модуля процесса с рангом 0 всем остальным репликам системы.
Примечание
Если вы используете DistributedDataParallel вместе с распределённой RPC-инфраструктурой, для вычисления градиентов всегда используйте
torch.distributed.autograd.backward(), а для оптимизации параметров —torch.distributed.optim.DistributedOptimizer.Пример:
>>> import torch.distributed.autograd as dist_autograd >>> from torch.nn.parallel import DistributedDataParallel as DDP >>> import torch >>> from torch import optim >>> from torch.distributed.optim import DistributedOptimizer >>> import torch.distributed.rpc as rpc >>> from torch.distributed.rpc import RRef >>> >>> t1 = torch.rand((3, 3), requires_grad=True) >>> t2 = torch.rand((3, 3), requires_grad=True) >>> rref = rpc.remote("worker1", torch.add, args=(t1, t2)) >>> ddp_model = DDP(my_model) >>> >>> # Setup optimizer >>> optimizer_params = [rref] >>> for param in ddp_model.parameters(): >>> optimizer_params.append(RRef(param)) >>> >>> dist_optim = DistributedOptimizer( >>> optim.SGD, >>> optimizer_params, >>> lr=0.05, >>> ) >>> >>> with dist_autograd.context() as context_id: >>> pred = ddp_model(rref.to_here()) >>> loss = loss_func(pred, target) >>> dist_autograd.backward(context_id, [loss]) >>> dist_optim.step(context_id)Примечание
В настоящее время DistributedDataParallel ограниченно поддерживает контрольные точки градиентов с помощью
torch.utils.checkpoint(). Если контрольная точка создаётся с use_reentrant=False (рекомендуется), DDP работает как ожидается без каких-либо ограничений. Если же контрольная точка создаётся с use_reentrant=True (значение по умолчанию), DDP работает как ожидается, если в модели нет неиспользуемых параметров и каждый слой сохраняется в контрольную точку не более одного раза (не передавайтеfind_unused_parameters=Trueв DDP). В настоящее время не поддерживается ситуация, когда слой сохраняется в контрольную точку несколько раз или в модели с контрольными точками есть неиспользуемые параметры.Примечание
Чтобы модель без DDP могла загрузить state dict модели с DDP, перед загрузкой необходимо применить
consume_prefix_in_state_dict_if_present(), чтобы удалить префикс «module.» из state dict DDP.Предупреждение
Конструктор, метод forward и дифференцирование выходных данных (или функции от выходных данных этого модуля) являются точками распределённой синхронизации. Учитывайте это, если в разных процессах может выполняться разный код.
Предупреждение
Этот модуль предполагает, что все параметры зарегистрированы в модели к моменту его создания. Позднее добавлять или удалять параметры нельзя. То же относится к буферам.
Предупреждение
Этот модуль предполагает, что параметры в моделях всех распределённых процессов зарегистрированы в одном и том же порядке. Сам модуль выполняет
allreduceградиентов в порядке, обратном порядку регистрации параметров модели. Иными словами, пользователь должен обеспечить, чтобы во всех распределённых процессах использовалась совершенно одинаковая модель и, следовательно, одинаковый порядок регистрации параметров.Предупреждение
Этот модуль допускает параметры с шагами, не соответствующими непрерывному расположению в порядке строк. Например, некоторые параметры модели могут иметь формат памяти
torch.memory_formattorch.contiguous_format, а другие — форматtorch.channels_last. Однако соответствующие параметры в разных процессах должны иметь одинаковые шаги.Предупреждение
Этот модуль несовместим с
torch.autograd.grad()(то есть он работает только в том случае, если градиенты накапливаются в атрибутах параметров.grad).Предупреждение
Если вы планируете использовать этот модуль с backend
ncclили backendgloo(использующим Infiniband) вместе с DataLoader, который задействует несколько рабочих процессов, измените метод запуска многопроцессной обработки наforkserverилиspawn. К сожалению, Gloo (использующий Infiniband) и NCCL2 небезопасны при разветвлении процесса (fork), поэтому без изменения этой настройки, вероятно, возникнут взаимоблокировки.Предупреждение
Никогда не пытайтесь изменять параметры модели после её обёртывания в
DistributedDataParallel. При обёртывании модели вDistributedDataParallelконструкторDistributedDataParallelрегистрирует дополнительные функции сведения градиентов для всех параметров модели, существующих на момент создания. Если изменить параметры модели после этого, функции сведения градиентов перестанут соответствовать нужному набору параметров.Предупреждение
Использование
DistributedDataParallelвместе с распределённой RPC-инфраструктурой является экспериментальным и может измениться.- Параметры:
-
- module (Module) – модуль, который необходимо распараллелить
-
device_ids (list из int или torch.device) –
Устройства CUDA. 1) Для модулей с одним устройством
device_idsможет содержать ровно один идентификатор устройства, обозначающий единственное устройство CUDA, на котором находится входной модуль данного процесса. В качестве альтернативыdevice_idsможет быть равенNone. 2) Для модулей с несколькими устройствами и модулей CPUdevice_idsдолжен быть равенNone.Если в обоих случаях
device_idsравенNone, входные данные для прямого прохода и сам модуль должны находиться на правильном устройстве. (по умолчанию:None) -
output_device (int или torch.device) – Устройство, на котором располагаются выходные данные модулей CUDA с одним устройством. Для модулей с несколькими устройствами и модулей CPU значение должно быть
None; расположение выходных данных определяется самим модулем. (по умолчанию:device_ids[0]для модулей с одним устройством) -
broadcast_buffers (bool или None) –
Флаг, включающий синхронизацию (рассылку) буферов модуля в начале функции
forward. (по умолчанию:None)Устарело начиная с версии 2.13: Вместо этого используйте
forward_sync_buffers. -
init_sync (bool) – Определяет, выполнять ли синхронизацию при инициализации для проверки формы параметров и рассылки параметров и буферов. Примечание: устаревший параметр
broadcast_buffers=Falseисключает буферы из этой начальной синхронизации. Заменаforward_sync_buffersне влияет на начальную синхронизацию. ПРЕДУПРЕЖДЕНИЕ: если задано False, пользователь должен самостоятельно обеспечить одинаковые веса на всех рангах. (по умолчанию:True) -
process_group – Группа процессов, используемая для распределённой операции all-reduce данных. Если задано
None, будет использована группа процессов по умолчанию, созданная с помощьюtorch.distributed.init_process_group(). (по умолчанию:None) -
bucket_cap_mb –
DistributedDataParallelраспределяет параметры по нескольким корзинам, чтобы сведение градиентов каждой корзины могло выполняться параллельно с обратным проходом.bucket_cap_mbзадаёт размер корзины в мебибайтах (МиБ). Если заданоNone, используется размер по умолчанию 25 МиБ. (по умолчанию:None) -
find_unused_parameters (bool) – Выполняет обход графа autograd, начиная со всех тензоров, содержащихся в значении, возвращаемом функцией
forwardобёрнутого модуля. Параметры, не получающие градиенты в этом графе, заранее помечаются как готовые к сведению. Кроме того, параметры, которые могли использоваться в функцииforwardобёрнутого модуля, но не участвовали в вычислении функции потерь и поэтому также не получат градиенты, заранее помечаются как готовые к сведению. (по умолчанию:False) - check_reduction – Этот аргумент устарел.
-
gradient_as_bucket_view (bool) – Если задано
True, градиенты становятся представлениями, указывающими на разные смещения в коммуникационных корзинахallreduce. Это может снизить пиковое потребление памяти; объём сэкономленной памяти будет равен общему размеру градиентов. Кроме того, это позволяет избежать затрат на копирование между градиентами и коммуникационными корзинамиallreduce. Когда градиенты являются представлениями, вызывать для нихdetach_()нельзя. Если возникает такая ошибка, воспользуйтесь функциейzero_grad()вtorch/optim/optimizer.py. Обратите внимание, что градиенты становятся представлениями после первой итерации, поэтому экономию пиковой памяти следует проверять после первой итерации. -
static_graph (bool) –
Если задано
True, DDP считает, что граф обучения статичен. Статический граф означает: 1) Набор используемых и неиспользуемых параметров не меняется на протяжении всего цикла обучения; в этом случае не имеет значения, задан ли параметрfind_unused_parameters = True. 2) Способ обучения графа не меняется на протяжении всего цикла обучения (то есть поток управления не зависит от итераций). Если для static_graph заданоTrue, DDP сможет обрабатывать случаи, которые ранее не поддерживались: 1) Реентерабельные обратные проходы. 2) Многократное создание контрольных точек активаций. 3) Создание контрольных точек активаций, если в модели есть неиспользуемые параметры. 4) Параметры модели, расположенные вне функции forward. 5) Потенциальное повышение производительности при наличии неиспользуемых параметров, поскольку при значенииTrueдля static_graph DDP не будет искать неиспользуемые параметры в графе на каждой итерации. Чтобы проверить, можно ли задать static_graph значениеTrue, можно посмотреть данные журнала DDP в конце предыдущего обучения модели. Еслиddp_logging_data.get("can_set_static_graph") == True, скорее всего, можно также задатьstatic_graph = True.- Пример::
-
>>> model_DDP = torch.nn.parallel.DistributedDataParallel(model) >>> # Training loop >>> ... >>> ddp_logging_data = model_DDP._get_ddp_logging_data() >>> static_graph = ddp_logging_data.get("can_set_static_graph")
-
delay_all_reduce_named_params (list из tuple из str и torch.nn.Parameter) – список именованных параметров, для которых операция all-reduce будет отложена до готовности градиента параметра, указанного в
param_to_hook_all_reduce. Другие аргументы DDP не применяются к именованным параметрам, указанным в этом аргументе, поскольку редьюсер DDP их игнорирует. -
param_to_hook_all_reduce (torch.nn.Parameter) – параметр, используемый для запуска отложенной операции all-reduce параметров, указанных в
delay_all_reduce_named_params. - skip_all_reduce_unused_params – Если задано значение True, DDP пропускает сведение неиспользуемых параметров. Для этого необходимо, чтобы неиспользуемые параметры оставались одинаковыми на всех рангах на протяжении всего процесса обучения. Если это условие не выполнено, возможны рассинхронизация и зависание обучения.
-
batched_grad_copy (bool) – Если задано
True, операции копирования градиентов отдельных параметров в корзины и деления откладываются и выполняются вместе как одна операция_foreach_copy_плюс одна плоская операцияdiv_при готовности корзины. Это сокращает число запусков ядер для каждого параметра до двух ядер на корзину и может повысить пропускную способность моделей со множеством небольших параметров. Оптимизация наиболее эффективна сoptimizer.zero_grad(set_to_none=True)(значение по умолчанию), когда один толькоgradient_as_bucket_viewне позволяет избежать копирования, поскольку псевдоним представления корзины уничтожается на каждой итерации. (по умолчанию:False) -
forward_sync_buffers (bool или None) – Флаг, включающий синхронизацию (рассылку) буферов модуля во время выполнения, в том числе в начале
forwardи после объединения процессов с неравным количеством входных данных. Не влияет на синхронизацию при инициализации (см.init_sync); буферы всегда синхронизируются при инициализации независимо от этого флага. Заменяет устаревший аргументbroadcast_buffers. Если заданоNone, по умолчанию используетсяTrue. (по умолчанию:None)
- Переменные:
-
module (Module) – модуль, который необходимо распараллелить.
Пример:
>>> torch.distributed.init_process_group(backend='nccl', world_size=4, init_method='...') >>> net = torch.nn.parallel.DistributedDataParallel(model)
-
join(divide_by_initial_world_size=True, enable=True, throw_on_early_termination=False)[исходный код] -
Менеджер контекста для обучения с неравным количеством входных данных в процессах DDP.
Этот менеджер контекста отслеживает уже завершившие работу процессы DDP и «имитирует» прямые и обратные проходы, добавляя коллективные коммуникационные операции, соответствующие операциям, создаваемым ещё работающими процессами DDP. Это гарантирует, что каждому коллективному вызову соответствует вызов уже завершивших работу процессов DDP, предотвращая зависания и ошибки, которые иначе возникли бы при обучении с неравным количеством входных данных в процессах. В качестве альтернативы, если для флага
throw_on_early_terminationзаданоTrue, все обучающие процессы вызовут ошибку, когда у одного из рангов закончатся входные данные; это позволит перехватить ошибки и обработать их в соответствии с логикой приложения.После присоединения всех процессов DDP менеджер контекста передаст модель последнего присоединившегося процесса всем процессам, чтобы убедиться, что модель во всех процессах одинакова (что гарантируется DDP).
Чтобы включить обучение с неравным количеством входных данных в процессах, просто оберните цикл обучения этим менеджером контекста. Никаких дополнительных изменений модели или загрузки данных не требуется.
Предупреждение
Если модель или цикл обучения, обёрнутые этим менеджером контекста, содержат дополнительные распределённые коллективные операции, например
SyncBatchNormв прямом проходе модели, необходимо включить флагthrow_on_early_termination. Это связано с тем, что менеджер контекста не учитывает коллективные коммуникации, не относящиеся к DDP. При включённом флаге все ранги вызовут исключение, когда у любого из них закончатся входные данные, что позволит перехватить и обработать эти ошибки на всех рангах.- Параметры:
-
-
divide_by_initial_world_size (bool) – Если задано
True, градиенты будут делиться на исходный размерworld_size, с которым было запущено обучение DDP. Если заданоFalse, будет вычислен фактический размер мира (количество рангов, у которых ещё не закончились входные данные), и градиенты при all-reduce будут делиться на него. Задайтеdivide_by_initial_world_size=True, чтобы каждый пример, включая примеры из неравномерных входных данных, вносил одинаковый вклад в глобальный градиент. Для этого градиент всегда делится на исходныйworld_size, даже при обработке неравномерных входных данных. Если задатьFalse, градиент будет делиться на число оставшихся узлов. Это обеспечивает соответствие обучению на меньшемworld_size, хотя и означает, что неравномерные входные данные будут вносить больший вклад в глобальный градиент. Обычно для случаев, когда неравномерными являются последние несколько входных данных задания обучения, следует задаватьTrue. В крайних случаях, при большом расхождении в количестве входных данных, значениеFalseможет дать лучшие результаты. -
enable (bool) – Включает или отключает обнаружение неравномерных входных данных. Передайте
enable=False, если известно, что количество входных данных одинаково во всех участвующих процессах. По умолчанию —True. -
throw_on_early_termination (bool) – Определяет, вызывать ли ошибку или продолжать обучение, когда хотя бы у одного ранга заканчиваются входные данные. Если задано
True, ошибка будет вызвана, когда первый ранг достигнет конца данных. Если заданоFalse, обучение продолжится с уменьшенным фактическим размером мира до присоединения всех рангов. Обратите внимание: если указан этот флаг, флагdivide_by_initial_world_sizeбудет проигнорирован. По умолчанию —False.
-
divide_by_initial_world_size (bool) – Если задано
Пример:
>>> import torch >>> import torch.distributed as dist >>> import os >>> import torch.multiprocessing as mp >>> import torch.nn as nn >>> # On each spawned worker >>> def worker(rank): >>> dist.init_process_group("nccl", rank=rank, world_size=2) >>> torch.cuda.set_device(rank) >>> model = nn.Linear(1, 1, bias=False).to(rank) >>> model = torch.nn.parallel.DistributedDataParallel( >>> model, device_ids=[rank], output_device=rank >>> ) >>> # Rank 1 gets one more input than rank 0. >>> inputs = [torch.tensor([1]).float() for _ in range(10 + rank)] >>> with model.join(): >>> for _ in range(5): >>> for inp in inputs: >>> loss = model(inp).sum() >>> loss.backward() >>> # Without the join() API, the below synchronization will hang >>> # blocking for rank 1's allreduce to complete. >>> torch.cuda.synchronize(device=rank)
-
join_hook(**kwargs)[исходный код] -
Хук присоединения DDP обеспечивает обучение с неравномерными входными данными, имитируя коммуникации в прямых и обратных проходах.
- Параметры:
-
kwargs (dict) –
dictс любыми именованными аргументами для изменения поведения хука присоединения во время выполнения; всем экземплярамJoinable, использующим один и тот же менеджер контекста присоединения, передаётся одинаковое значениеkwargs.
- Хук поддерживает следующие именованные аргументы:
-
- divide_by_initial_world_size (bool, optional):
-
Если задано
True, градиенты делятся на исходный размер мира, с которым было запущено DDP. Если заданоFalse, градиенты делятся на фактический размер мира (то есть количество ещё не присоединившихся процессов), поэтому неравномерные входные данные вносят больший вклад в глобальный градиент. Обычно следует задатьTrue, если степень неравномерности невелика; в крайних случаях для получения потенциально лучших результатов можно задатьFalse. По умолчанию —True.
-
no_sync()[исходный код] -
Менеджер контекста для отключения синхронизации градиентов между процессами DDP.
Внутри этого контекста градиенты накапливаются в переменных модуля и синхронизируются при первом прямом и обратном проходе после выхода из контекста.
Пример:
>>> ddp = torch.nn.parallel.DistributedDataParallel(model, pg) >>> with ddp.no_sync(): >>> for input in inputs: >>> ddp(input).backward() # no synchronization, accumulate grads >>> ddp(another_input).backward() # synchronize grads
Предупреждение
Прямой проход необходимо выполнять внутри менеджера контекста; в противном случае градиенты всё равно будут синхронизироваться.
-
register_comm_hook(state, hook)[source] -
Зарегистрировать хук коммуникации для пользовательской агрегации градиентов DDP между несколькими рабочими процессами.
Этот хук может быть очень полезен исследователям для проверки новых идей. Например, его можно использовать для реализации таких алгоритмов, как GossipGrad и сжатие градиентов, предполагающих различные стратегии обмена данными для синхронизации параметров при обучении с использованием Distributed DataParallel.
- Параметры:
-
-
state (object) –
Передается хуку для хранения любой информации о состоянии в процессе обучения. Например, это могут быть данные обратной связи об ошибке при сжатии градиентов, список узлов для следующего обмена данными в GossipGrad и т. д.
Состояние хранится локально в каждом рабочем процессе и является общим для всех тензоров градиентов в этом процессе.
-
hook (Callable) –
Вызываемый объект со следующей сигнатурой:
hook(state: object, bucket: dist.GradBucket) -> torch.futures.Future[torch.Tensor]:Эта функция вызывается, когда корзина готова. Хук может выполнить любую необходимую обработку и вернуть Future, указывающий на завершение асинхронной работы (например, allreduce). Если хук не выполняет обмен данными, он все равно должен вернуть завершенный Future. Future должен содержать новое значение тензоров корзины градиентов. Когда корзина готова, редуктор c10d вызывает этот хук, использует тензоры, возвращенные Future, и копирует градиенты в отдельные параметры. Обратите внимание, что тип возвращаемого значения Future должен быть одним тензором.
Также мы предоставляем API под названием
get_futureдля получения Future, связанного с завершениемc10d.ProcessGroup.Work.get_futureв настоящее время поддерживается для NCCL, а также для большинства операций в GLOO и MPI, за исключением операций обмена данными между узлами (send/recv).
-
Предупреждение
Тензоры корзины градиентов не будут предварительно делиться на world_size. В случае таких операций, как allreduce, пользователь должен разделить их на world_size.
Предупреждение
Хук коммуникации DDP можно зарегистрировать только один раз, и его следует зарегистрировать до вызова backward.
Предупреждение
Объект Future, возвращаемый хуком, должен содержать один тензор той же формы, что и тензоры в корзине градиентов.
Предупреждение
API
get_futureподдерживает NCCL и частично бэкенды GLOO и MPI (не поддерживаются операции обмена данными между узлами, такие как send/recv) и возвращаетtorch.futures.Future.- Пример::
-
Ниже приведен пример пустого хука, который возвращает тот же тензор.
>>> def noop(state: object, bucket: dist.GradBucket) -> torch.futures.Future[torch.Tensor]: >>> fut = torch.futures.Future() >>> fut.set_result(bucket.buffer()) >>> return fut >>> ddp.register_comm_hook(state=None, hook=noop)
- Пример::
-
Ниже приведен пример алгоритма параллельного SGD, в котором градиенты кодируются перед allreduce, а затем декодируются после allreduce.
>>> def encode_and_decode(state: object, bucket: dist.GradBucket) -> torch.futures.Future[torch.Tensor]: >>> encoded_tensor = encode(bucket.buffer()) # encode gradients >>> fut = torch.distributed.all_reduce(encoded_tensor).get_future() >>> # Define the then callback to decode. >>> def decode(fut): >>> decoded_tensor = decode(fut.value()[0]) # decode gradients >>> return decoded_tensor >>> return fut.then(decode) >>> ddp.register_comm_hook(state=None, hook=encode_and_decode)
-
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.parallel.DistributedDataParallel.html