Spec-Zone.ru › PyTorch 2.14

DataParallel

class torch.nn.DataParallel(module, device_ids=None, output_device=None, dim=0) [исходный код]

Реализует параллелизм данных на уровне модуля.

Этот контейнер распараллеливает применение заданного module, разделяя входные данные между указанными устройствами по частям вдоль размерности пакета (другие объекты будут скопированы для каждого устройства). При прямом проходе модуль реплицируется на каждом устройстве, и каждая реплика обрабатывает часть входных данных. Во время обратного прохода градиенты от каждой реплики суммируются в исходном модуле.

Размер пакета должен быть больше числа используемых GPU.

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

Для обучения на нескольких GPU рекомендуется использовать DistributedDataParallel вместо этого класса, даже если используется только один узел. См. разделы: Используйте nn.parallel.DistributedDataParallel вместо multiprocessing или nn.DataParallel и Распределённый параллелизм данных.

В DataParallel можно передавать произвольные позиционные и именованные входные аргументы, однако некоторые типы обрабатываются особым образом. Тензоры будут распределены по указанной размерности (по умолчанию 0). Объекты типов tuple, list и dict будут неглубоко скопированы. Объекты других типов будут совместно использоваться разными потоками и могут быть повреждены, если модель изменит их во время прямого прохода.

Параллелизуемый module должен иметь параметры и буферы на device_ids[0] перед запуском этого модуля DataParallel.

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

При каждом прямом проходе module реплицируется на каждом устройстве, поэтому любые обновления работающего модуля в forward будут потеряны. Например, если у module есть атрибут-счётчик, который увеличивается при каждом forward, его значение всегда будет оставаться исходным, поскольку обновление выполняется на репликах, которые уничтожаются после forward. Однако DataParallel гарантирует, что реплика на device[0] будет совместно использовать хранилище параметров и буферов с базовым параллелизуемым module. Поэтому изменения на месте параметров или буферов на device[0] будут сохранены. Например, BatchNorm2d и spectral_norm() полагаются на это поведение для обновления буферов.

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

Обработчики прямого и обратного проходов, определённые для module и его подмодулей, будут вызваны len(device_ids) раз, каждый раз с входными данными, расположенными на определённом устройстве. В частности, гарантируется только правильный порядок выполнения обработчиков относительно операций на соответствующих устройствах. Например, не гарантируется, что обработчики, заданные с помощью register_forward_pre_hook(), выполнятся до вызовов all len(device_ids) forward(), но гарантируется, что каждый такой обработчик выполнится до соответствующего вызова forward() на данном устройстве.

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

Если module возвращает скаляр (то есть тензор нулевой размерности) в forward(), эта обёртка вернёт вектор длиной, равной числу устройств, используемых для параллелизма данных, содержащий результат с каждого устройства.

Примечание

При использовании шаблона pack sequence -> recurrent network -> unpack sequence в Module, обёрнутом в DataParallel, есть один нюанс. Подробности см. в разделе часто задаваемых вопросов Моя рекуррентная сеть не работает с параллелизмом данных.

Параметры:
  • module (Module) – модуль, который нужно распараллелить
  • device_ids (list of int or torch.device) – устройства CUDA (по умолчанию: все устройства)
  • output_device (int or torch.device) – устройство для размещения выходных данных (по умолчанию: device_ids[0])
Переменные:

module (Module) – модуль, который нужно распараллелить

Пример:

>>> net = torch.nn.DataParallel(model, device_ids=[0, 1, 2])
>>> output = net(input_var)  # input_var can be on any device, including CPU

© 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.DataParallel.html

Spec-Zone.ru

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