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