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(), выполнятся до вызововalllen(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