Spec-Zone.ru › PyTorch 2.14

Распределённые оптимизаторы

Создано: 01 мар. 2021 | Последнее обновление: 11 мая 2026

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

Распределённый оптимизатор в настоящее время не поддерживается при использовании тензоров CUDA

torch.distributed.optim предоставляет DistributedOptimizer, который принимает список удалённых параметров (RRef) и запускает оптимизатор локально на рабочих узлах, где находятся параметры. Распределённый оптимизатор может использовать любой локальный базовый класс оптимизатора для применения градиентов на каждом рабочем узле.

class torch.distributed.optim.DistributedOptimizer(optimizer_class, params_rref, *args, **kwargs) [исходный код]

DistributedOptimizer принимает удалённые ссылки на параметры, распределённые между рабочими узлами, и применяет заданный оптимизатор локально для каждого параметра.

Этот класс использует get_gradients() для получения градиентов определённых параметров.

Параллельные вызовы step() — как от одного, так и от разных клиентов — будут сериализованы на каждом рабочем узле, поскольку оптимизатор каждого узла может работать только с одним набором градиентов одновременно. Однако нет гарантии, что полная последовательность прямого прохода, обратного прохода и работы оптимизатора будет выполняться для одного клиента за раз. Это означает, что применяемые градиенты могут не соответствовать последнему прямому проходу, выполненному на данном рабочем узле. Кроме того, порядок выполнения на разных рабочих узлах не гарантируется.

DistributedOptimizer использует функциональный оптимизатор, если он доступен для заданного optimizer_class, чтобы обновления оптимизатора не блокировались глобальной блокировкой интерпретатора Python (GIL) при многопоточной обучении (например, при распределённом параллелизме модели). В настоящее время эта возможность включена для большинства оптимизаторов.

Параметры:
  • optimizer_class (optim.Optimizer) – класс оптимизатора, который нужно создать на каждом рабочем узле.
  • params_rref (list[RRef]) – список RRef для локальных или удалённых параметров, подлежащих оптимизации.
  • args – аргументы, передаваемые конструктору оптимизатора на каждом рабочем узле.
  • kwargs – аргументы, передаваемые конструктору оптимизатора на каждом рабочем узле.
Пример::
>>> import torch.distributed.autograd as dist_autograd
>>> import torch.distributed.rpc as rpc
>>> from torch import optim
>>> from torch.distributed.optim import DistributedOptimizer
>>>
>>> with dist_autograd.context() as context_id:
>>>   # Forward pass.
>>>   rref1 = rpc.remote("worker1", torch.add, args=(torch.ones(2), 3))
>>>   rref2 = rpc.remote("worker1", torch.add, args=(torch.ones(2), 1))
>>>   loss = rref1.to_here() + rref2.to_here()
>>>
>>>   # Backward pass.
>>>   dist_autograd.backward(context_id, [loss.sum()])
>>>
>>>   # Optimizer.
>>>   dist_optim = DistributedOptimizer(
>>>      optim.SGD,
>>>      [rref1, rref2],
>>>      lr=0.05,
>>>   )
>>>   dist_optim.step(context_id)
step(context_id) [исходный код]

Выполняет один шаг оптимизации.

Этот метод вызовет torch.optim.Optimizer.step() на каждом рабочем узле, содержащем параметры для оптимизации, и будет блокировать выполнение, пока не ответят все рабочие узлы. Указанный context_id будет использоваться для получения соответствующего context, содержащего градиенты, которые следует применить к параметрам.

Параметры:

context_id – идентификатор контекста автоградиента, для которого следует выполнить шаг оптимизатора.

class torch.distributed.optim.PostLocalSGDOptimizer(optim, averager) [исходный код]

Оборачивает произвольный torch.optim.Optimizer и выполняет post-local SGD. Этот оптимизатор запускает локальный оптимизатор на каждом шаге. После этапа разогрева он периодически усредняет параметры после применения локального оптимизатора.

Параметры:
  • optim (Optimizer) – локальный оптимизатор.
  • averager (ModelAverager) – экземпляр средства усреднения модели для выполнения алгоритма post-localSGD.

Пример:

>>> import torch
>>> import torch.distributed as dist
>>> import torch.distributed.algorithms.model_averaging.averagers as averagers
>>> import torch.nn as nn
>>> from torch.distributed.optim import PostLocalSGDOptimizer
>>> from torch.distributed.algorithms.ddp_comm_hooks.post_localSGD_hook import (
>>>   PostLocalSGDState,
>>>   post_localSGD_hook,
>>> )
>>>
>>> model = nn.parallel.DistributedDataParallel(
>>>    module, device_ids=[rank], output_device=rank
>>> )
>>>
>>> # Register a post-localSGD communication hook.
>>> state = PostLocalSGDState(process_group=None, subgroup=None, start_localSGD_iter=100)
>>> model.register_comm_hook(state, post_localSGD_hook)
>>>
>>> # Create a post-localSGD optimizer that wraps a local optimizer.
>>> # Note that ``warmup_steps`` used in ``PostLocalSGDOptimizer`` must be the same as
>>> # ``start_localSGD_iter`` used in ``PostLocalSGDState``.
>>> local_optim = torch.optim.SGD(params=model.parameters(), lr=0.01)
>>> opt = PostLocalSGDOptimizer(
>>>     optim=local_optim,
>>>     averager=averagers.PeriodicModelAverager(period=4, warmup_steps=100)
>>> )
>>>
>>> # In the first 100 steps, DDP runs global gradient averaging at every step.
>>> # After 100 steps, DDP runs gradient averaging within each subgroup (intra-node by default),
>>> # and post-localSGD optimizer runs global model averaging every 4 steps after applying the local optimizer.
>>> for step in range(0, 200):
>>>    opt.zero_grad()
>>>    loss = loss_fn(output, labels)
>>>    loss.backward()
>>>    opt.step()
load_state_dict(state_dict) [исходный код]

Работает так же, как torch.optim.Optimizer load_state_dict(), но также восстанавливает значение шага средства усреднения модели из предоставленного state_dict.

Если в state_dict отсутствует запись "step", будет выдано предупреждение, а значение шага средства усреднения модели будет инициализировано нулём.

state_dict() [исходный код]

Работает так же, как torch.optim.Optimizer state_dict(), но добавляет в контрольную точку дополнительную запись со значением шага средства усреднения модели, чтобы при повторной загрузке не происходил ненужный повторный этап разогрева.

step() [исходный код]

Выполняет один шаг оптимизации (обновление параметров).

class torch.distributed.optim.ZeroRedundancyOptimizer(params, optimizer_class, process_group=None, parameters_as_bucket_view=False, overlap_with_ddp=False, **defaults) [исходный код]

Оборачивает произвольный optim.Optimizer и распределяет его состояния между рангами группы.

Совместное использование выполняется согласно описанию в статье ZeRO.

Локальный экземпляр оптимизатора на каждом ранге отвечает только за обновление примерно 1 / world_size параметров, поэтому ему нужно хранить только 1 / world_size состояний оптимизатора. После локального обновления параметров каждый ранг рассылает свои параметры всем остальным узлам, чтобы все реплики модели находились в одинаковом состоянии. ZeroRedundancyOptimizer можно использовать совместно с torch.nn.parallel.DistributedDataParallel, чтобы уменьшить пиковое потребление памяти на каждом ранге.

ZeroRedundancyOptimizer использует алгоритм жадной упаковки с предварительной сортировкой для распределения параметров между рангами. Каждый параметр принадлежит одному рангу и не разделяется между рангами. Разбиение произвольно и может не совпадать с порядком регистрации или использования параметров.

Параметры:

params (Iterable) – Iterable из torch.Tensor или dict, содержащий все параметры, которые будут распределены между рангами.

Именованные аргументы:
  • optimizer_class (torch.nn.Optimizer) – класс локального оптимизатора.
  • process_group (ProcessGroup, необязательно) – torch.distributed ProcessGroup (по умолчанию: dist.group.WORLD, инициализированная с помощью torch.distributed.init_process_group()).
  • parameters_as_bucket_view (bool, необязательно) – если True, параметры упаковываются в бакеты для ускорения обмена данными, а поля param.data указывают на представления бакетов со смещениями; если False, каждый параметр передаётся отдельно, и каждый params.data остаётся неизменным (по умолчанию: False).
  • overlap_with_ddp (bool, необязательно) – если True, step() выполняется параллельно с синхронизацией градиентов в DistributedDataParallel; для этого требуется (1) функциональный оптимизатор для аргумента optimizer_class или оптимизатор с функциональным эквивалентом и (2) регистрация обработчика обмена данными DDP, созданного с помощью одной из функций в ddp_zero_hook.py; параметры упаковываются в бакеты, соответствующие бакетам в DistributedDataParallel, поэтому аргумент parameters_as_bucket_view игнорируется. Если False, step() выполняется отдельно после обратного прохода (как обычно). (по умолчанию: False)
  • **defaults – любые дополнительные конечные аргументы, передаваемые локальному оптимизатору.

Пример:

>>> import torch.nn as nn
>>> from torch.distributed.optim import ZeroRedundancyOptimizer
>>> from torch.nn.parallel import DistributedDataParallel as DDP
>>> model = nn.Sequential(*[nn.Linear(2000, 2000).to(rank) for _ in range(20)])
>>> ddp = DDP(model, device_ids=[rank])
>>> opt = ZeroRedundancyOptimizer(
>>>     ddp.parameters(),
>>>     optimizer_class=torch.optim.Adam,
>>>     lr=0.01
>>> )
>>> ddp(inputs).sum().backward()
>>> opt.step()

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

В настоящее время ZeroRedundancyOptimizer требует, чтобы все переданные параметры имели один и тот же плотный тип.

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

При передаче overlap_with_ddp=True учитывайте следующее: из-за текущей реализации параллельного выполнения DistributedDataParallel с ZeroRedundancyOptimizer первые две или три итерации обучения не выполняют обновления параметров на шаге оптимизатора — в зависимости от того, используется ли static_graph=False или static_graph=True соответственно. Это происходит потому, что необходима информация о стратегии разбиения градиентов на бакеты, используемой DistributedDataParallel, которая окончательно определяется только после второго прямого прохода при static_graph=False или после третьего прямого прохода при static_graph=True. Чтобы учесть это, можно добавить фиктивные входные данные в начало.

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

ZeroRedundancyOptimizer является экспериментальным и может измениться.

add_param_group(param_group) [исходный код]

Добавляет группу параметров в param_groups объекта Optimizer.

Это может быть полезно при тонкой настройке предварительно обученной сети: замороженные слои можно сделать обучаемыми и добавлять в Optimizer по мере обучения.

Параметры:

param_group (dict) – задаёт параметры для оптимизации и специфичные для группы параметры оптимизации.

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

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

consolidate_state_dict(to=0) [исходный код]

Собирает список state_dict (по одному на ранг) на целевом ранге.

Параметры:

to (int) – ранг, который получает состояния оптимизатора (по умолчанию: 0).

Исключения:

RuntimeError – если overlap_with_ddp=True и этот метод вызван до полной инициализации экземпляра ZeroRedundancyOptimizer, которая происходит после перестроения бакетов градиентов DistributedDataParallel.

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

Этот метод необходимо вызывать на всех рангах.

property join_device: device

Возвращает устройство по умолчанию.

join_hook(**_kwargs) [исходный код]

Возвращает хук присоединения ZeRO.

Он позволяет обучаться на неодинаковых входных данных, имитируя коллективный обмен данными на шаге оптимизатора.

Перед вызовом этого хука необходимо правильно задать градиенты.

Параметры:

kwargs (dict) – dict, содержащий любые именованные аргументы для изменения поведения хука присоединения во время выполнения; всем экземплярам Joinable, использующим один и тот же менеджер контекста присоединения, передаётся одно и то же значение kwargs.

Тип возвращаемого значения:

JoinHook

Этот хук не поддерживает именованные аргументы; то есть kwargs не используется.

property join_process_group: Any

Возвращает группу процессов.

load_state_dict(state_dict) [исходный код]

Загружает состояние, относящееся к данному рангу, из входного state_dict и при необходимости обновляет локальный оптимизатор.

Параметры:

state_dict (dict) – состояние оптимизатора; должен быть объектом, возвращённым вызовом state_dict().

Исключения:

RuntimeError – если overlap_with_ddp=True и этот метод вызван до полной инициализации экземпляра ZeroRedundancyOptimizer, которая происходит после перестроения бакетов градиентов DistributedDataParallel.

state_dict() [исходный код]

Возвращает последнее глобальное состояние оптимизатора, известное этому рангу.

Исключения:

RuntimeError – если overlap_with_ddp=True и этот метод вызван до полной инициализации экземпляра ZeroRedundancyOptimizer, которая происходит после перестроения бакетов градиентов DistributedDataParallel; либо если этот метод вызван без предварительного вызова consolidate_state_dict().

Тип возвращаемого значения:

dict[str, Any]

step(closure=None, **kwargs) [исходный код]

Выполняет один шаг оптимизатора и синхронизирует параметры между всеми рангами.

Параметры:

closure (Callable) – замыкание, повторно вычисляющее модель и возвращающее значение функции потерь; для большинства оптимизаторов необязательно.

Возвращает:

Необязательное значение функции потерь, зависящее от используемого локального оптимизатора.

Тип возвращаемого значения:

float | None

Примечание

Все дополнительные параметры передаются базовому оптимизатору без изменений.

torch.distributed.optim.utils.register_functional_optim(key, optim) [исходный код]

Интерфейс для добавления нового функционального оптимизатора в functional_optim_map. fn_optim_key и fn_optimizer задаются пользователем. Оптимизатор и ключ не обязаны быть экземплярами torch.optim.Optimizer (например, для пользовательских оптимизаторов).

Пример:

>>> # import the new functional optimizer
>>> from xyz import fn_optimizer
>>> from torch.distributed.optim.utils import register_functional_optim
>>> fn_optim_key = "XYZ_optim"
>>> register_functional_optim(fn_optim_key, fn_optimizer)

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/distributed.optim.html

Spec-Zone.ru

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