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(), будут выполнены доalllen(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