Spec-Zone.ru › PyTorch 1

DataParallel

class torch.nn.DataParallel(module, device_ids=None, output_device=None, dim=0) [source]

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

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

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

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

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

Разрешается передавать произвольные позиционные и ключевые аргументы в DataParallel, но некоторые типы обрабатываются особым образом. тензоры будут разбросаны по указанной размерности (по умолчанию 0). кортежи, списки и словари будут созданы как поверхностные копии. Другие типы будут совместно использоваться между различными потоками и могут быть повреждены, если к ним обращаться в прямом проходе модели.

Распараллеленный 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 возвращает скаляр (т.е. тензор размерности 0) в forward(), этот обёртка вернёт вектор длиной, равной количеству используемых устройств в распараллеливании данных, содержащий результат с каждого устройства.

Примечание

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

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

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

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.nn.DataParallel.html

Spec-Zone.ru

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