Spec-Zone.ru › PyTorch 2

DataParallel

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

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

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

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

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

Рекомендуется использовать DistributedDataParallel вместо этого класса для обучения на нескольких графических процессорах, даже если используется только один узел. См.: Использование 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.

Параметры
  • module (Module) – модуль, подлежащий параллелизации
  • device_ids (список из целых чисел или torch.device) – CUDA устройства (по умолчанию: все устройства)
  • output_device (целое число или 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

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

Spec-Zone.ru

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