Spec-Zone.ru › PyTorch 2.14

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_location torch.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_format torch.contiguous_format, а другие — формат torch.channels_last. Однако соответствующие параметры в разных процессах должны иметь одинаковые шаги.

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

Этот модуль несовместим с torch.autograd.grad() (то есть он работает только в том случае, если градиенты накапливаются в атрибутах параметров .grad).

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

Если вы планируете использовать этот модуль с backend nccl или backend gloo (использующим 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) Для модулей с несколькими устройствами и модулей CPU device_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.

Пример:

>>> 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

Spec-Zone.ru

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