Spec-Zone.ru › PyTorch 2.14

Архитектура распределённого автограда

Дата создания: 8 мая 2026 г. | Дата последнего обновления: 8 мая 2026 г.

В этой заметке подробно описана архитектура распределённого автограда и рассмотрены его внутренние механизмы. Перед продолжением ознакомьтесь с механикой автограда и распределённой RPC-средой.

Общие сведения

Предположим, у вас есть два узла и очень простая модель, разделённая между ними. Её можно реализовать с помощью torch.distributed.rpc следующим образом:

import torch
import torch.distributed.rpc as rpc

def my_add(t1, t2):
  return torch.add(t1, t2)

# On worker 0:
t1 = torch.rand((3, 3), requires_grad=True)
t2 = torch.rand((3, 3), requires_grad=True)

# Perform some computation remotely.
t3 = rpc.rpc_sync("worker1", my_add, args=(t1, t2))

# Perform some computation locally based on remote result.
t4 = torch.rand((3, 3), requires_grad=True)
t5 = torch.mul(t3, t4)

# Compute some loss.
loss = t5.sum()

Основная цель распределённого автограда — обеспечить выполнение обратного прохода для таких распределённых моделей с использованием loss, которые мы вычислили, и запись соответствующих градиентов для всех тензоров, которым требуются градиенты.

Запись автограда во время прямого прохода

PyTorch строит граф автограда во время прямого прохода, и этот граф используется для выполнения обратного прохода. Подробнее см. в разделе Как автоград кодирует историю вычислений.

Для распределённого автограда необходимо отслеживать все RPC во время прямого прохода, чтобы обеспечить корректное выполнение обратного прохода. Для этого при выполнении RPC мы добавляем в граф автограда функции send и recv.

  • Функция send добавляется в источник RPC, а её выходные рёбра указывают на функцию автограда для входных тензоров RPC. На вход этой функции во время обратного прохода поступает значение, полученное от узла назначения как результат соответствующей функции recv.
  • Функция recv добавляется в узел назначения, а её входные данные извлекаются из операций, выполненных на узле назначения с использованием входных тензоров. Выходные градиенты этой функции во время обратного прохода отправляются на исходный узел соответствующей функции send.
  • Каждой паре send-recv присваивается глобально уникальный autograd_message_id, позволяющий однозначно идентифицировать эту пару. Это удобно для поиска соответствующей функции на удалённом узле во время обратного прохода.
  • Для RRef каждый раз, когда мы вызываем torch.distributed.rpc.RRef.to_here(), для задействованных тензоров добавляется соответствующая пара send-recv.

Например, граф автограда для приведённого выше примера будет выглядеть так (для простоты t5.sum() не показан):

send and recv functions

Контекст распределённого автограда

Каждому прямому и обратному проходу, использующему распределённый автоград, назначается уникальный torch.distributed.autograd.context, у которого есть глобально уникальный autograd_context_id. Этот контекст создаётся на каждом узле по мере необходимости.

Этот контекст выполняет следующие задачи:

  1. Несколько узлов, выполняющих распределённые обратные проходы, могут накапливать градиенты для одного и того же тензора, поэтому поле .grad тензора будет содержать градиенты из различных распределённых обратных проходов, прежде чем мы успеем запустить оптимизатор. Это аналогично многократному локальному вызову torch.autograd.backward(). Чтобы разделять градиенты для каждого обратного прохода, они накапливаются в torch.distributed.autograd.context соответствующего прохода.
  2. Во время прямого прохода мы сохраняем в этом контексте функции send и recv для каждого прохода автограда. Это позволяет удерживать ссылки на соответствующие узлы графа автограда, чтобы он оставался доступным. Кроме того, это упрощает поиск нужных функций send и recv во время обратного прохода.
  3. В целом мы также используем этот контекст для хранения метаданных каждого распределённого прохода автограда.

С точки зрения пользователя контекст автограда настраивается следующим образом:

import torch.distributed.autograd as dist_autograd
with dist_autograd.context() as context_id:
  loss = model.forward()
  dist_autograd.backward(context_id, loss)

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

Распределённый обратный проход

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

Вычисление зависимостей

Рассмотрим следующий фрагмент кода, выполняемый на одном компьютере:

import torch
a = torch.rand((3, 3), requires_grad=True)
b = torch.rand((3, 3), requires_grad=True)
c = torch.rand((3, 3), requires_grad=True)
d = a + b
e = b * c
d.sum().backward()

Граф автограда для приведённого выше кода будет выглядеть так:

local dependencies

Первый шаг, который выполняет движок автограда в рамках обратного прохода, — вычисление количества зависимостей для каждого узла графа автограда. Это помогает движку определить, когда узел графа готов к выполнению. Числа в скобках для add(1) и mul(0) обозначают количество зависимостей. Как видите, во время обратного прохода узлу add требуется один вход, а узлу mul входы не нужны (другими словами, его не нужно выполнять). Локальный движок автограда вычисляет эти зависимости, обходя граф от корневых узлов (в данном случае d).

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

import torch
import torch.distributed.rpc as rpc

a = torch.rand((3, 3), requires_grad=True)
b = torch.rand((3, 3), requires_grad=True)
c = torch.rand((3, 3), requires_grad=True)

d = rpc.rpc_sync("worker1", torch.add, args=(a, b))
e = rpc.rpc_sync("worker1", torch.mul, args=(b, c))
loss = d.sum()

Соответствующий граф автограда для приведённого выше кода будет выглядеть так:

distributed dependencies

Вычисление зависимостей для этого распределённого графа автограда гораздо сложнее и требует дополнительных затрат (на вычисления или сетевое взаимодействие).

Для приложений, чувствительных к производительности, можно избежать значительных дополнительных затрат, предположив, что каждая функция send и recv участвует в обратном проходе (в большинстве приложений RPC, результаты которых не используются, не выполняются). Это упрощает алгоритм распределённого автограда и значительно повышает его эффективность, но требует от приложения учитывать его ограничения. Этот алгоритм называется алгоритмом режима FAST и подробно описан ниже.

В общем случае не обязательно, что каждая функция send и recv участвует в обратном проходе. Для решения этой задачи предложен алгоритм режима SMART, который описан в одном из следующих разделов. Обратите внимание, что в настоящее время реализован только алгоритм режима FAST.

Алгоритм режима FAST

Ключевое предположение этого алгоритма заключается в том, что для каждой функции send при выполнении обратного прохода имеется одна зависимость. Иными словами, мы предполагаем, что получим градиент по RPC с другого узла.

Алгоритм работает следующим образом:

  1. Начинаем с рабочего узла, на котором находятся корни обратного прохода (все корни должны быть локальными).
  2. Ищем все функции send для текущего контекста распределённого автограда.
  3. Вычисляем зависимости локально, начиная с указанных корней и всех полученных функций send.
  4. После вычисления зависимостей запускаем локальный движок автограда с указанными корнями.
  5. Когда движок автограда выполняет функцию recv, функция recv отправляет входные градиенты по RPC соответствующему рабочему узлу. Каждая функция recv знает идентификатор рабочего узла назначения, поскольку он записывается во время прямого прохода. Функция recv также отправляет удалённому узлу autograd_context_id и autograd_message_id.
  6. Получив этот запрос на удалённом узле, мы используем autograd_context_id и autograd_message_id для поиска соответствующей функции send.
  7. Если рабочий узел впервые получает запрос для указанного autograd_context_id, он локально вычисляет зависимости, как описано выше в пунктах 1–3.
  8. Затем полученная в пункте 6 функция send ставится в очередь на выполнение локальным движком автограда этого рабочего узла.
  9. Наконец, вместо накопления градиентов в поле .grad тензора мы накапливаем их отдельно для каждого контекста распределённого автограда. Градиенты хранятся в Dict[Tensor, Tensor] — по сути, это отображение тензоров в соответствующие градиенты. Получить это отображение можно с помощью API get_gradients().

Например, полный код с распределённым автоградом будет выглядеть так:

import torch
import torch.distributed.autograd as dist_autograd
import torch.distributed.rpc as rpc

def my_add(t1, t2):
  return torch.add(t1, t2)

# On worker 0:

# Setup the autograd context. Computations that take
# part in the distributed backward pass must be within
# the distributed autograd context manager.
with dist_autograd.context() as context_id:
  t1 = torch.rand((3, 3), requires_grad=True)
  t2 = torch.rand((3, 3), requires_grad=True)

  # Perform some computation remotely.
  t3 = rpc.rpc_sync("worker1", my_add, args=(t1, t2))

  # Perform some computation locally based on remote result.
  t4 = torch.rand((3, 3), requires_grad=True)
  t5 = torch.mul(t3, t4)

  # Compute some loss.
  loss = t5.sum()

  # Run the backward pass.
  dist_autograd.backward(context_id, [loss])

  # Retrieve the gradients from the context.
  dist_autograd.get_gradients(context_id)

Распределённый граф автограда с зависимостями будет выглядеть следующим образом (для простоты t5.sum() не показан):

distributed dependencies computed

Алгоритм режима FAST применительно к приведённому выше примеру будет работать следующим образом:

  1. На Worker 0 начинаем с корней loss и send1, чтобы вычислить зависимости. В результате для send1 отмечается одна зависимость, а для mul на Worker 0 также отмечается одна зависимость.
  2. Теперь запускаем локальный движок автограда на Worker 0. Сначала выполняем функцию mul и накапливаем её результат в контексте автограда как градиент для t4. Затем выполняем recv2, которая отправляет градиенты в Worker 1.
  3. Поскольку Worker 1 впервые получает сведения об этом обратном проходе, он начинает вычисление зависимостей и соответствующим образом отмечает зависимости для send2, add и recv1.
  4. Затем ставим send2 в очередь локального движка автограда на Worker 1, который, в свою очередь, выполняет add и recv1.
  5. При выполнении recv1 градиенты отправляются в Worker 0.
  6. Поскольку Worker 0 уже вычислил зависимости для этого обратного прохода, он просто ставит send1 в очередь и выполняет её локально.
  7. Наконец, градиенты для t1, t2 и t4 накапливаются в контексте распределённого автограда.

Алгоритм режима SMART

Работа над подробным описанием этого алгоритма ещё продолжается. Общее представление о нём можно получить из раздела Алгоритм распределённого автограда: режим Smart в документе RFC.

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

DistributedOptimizer работает следующим образом:

  1. Принимает список удалённых параметров (RRef) для оптимизации. Это также могут быть локальные параметры, обёрнутые в локальный RRef.
  2. Принимает класс Optimizer в качестве локального оптимизатора, который будет запущен для всех уникальных владельцев RRef.
  3. Распределённый оптимизатор создаёт экземпляр локального Optimizer на каждом рабочем узле и хранит RRef на эти экземпляры.
  4. При вызове torch.distributed.optim.DistributedOptimizer.step() распределённый оптимизатор использует RPC для удалённого запуска всех локальных оптимизаторов на соответствующих рабочих узлах. В качестве входных данных для torch.distributed.optim.DistributedOptimizer.step() необходимо передать контекст распределённого автограда context_id. Локальные оптимизаторы используют его для применения градиентов, сохранённых в соответствующем контексте.
  5. Если несколько параллельно работающих распределённых оптимизаторов обновляют одни и те же параметры на рабочем узле, эти обновления сериализуются с помощью блокировки.

Простой сквозной пример

Объединив всё вместе, получим простой сквозной пример использования распределённого автограда и распределённого оптимизатора. Если сохранить код в файл с именем dist_autograd_simple.py, его можно запустить командой MASTER_ADDR="localhost" MASTER_PORT=29500 python dist_autograd_simple.py:

import torch
import torch.multiprocessing as mp
import torch.distributed.autograd as dist_autograd
from torch.distributed import rpc
from torch import optim
from torch.distributed.optim import DistributedOptimizer

def random_tensor():
    return torch.rand((3, 3), requires_grad=True)

def _run_process(rank, dst_rank, world_size):
    name = "worker{}".format(rank)
    dst_name = "worker{}".format(dst_rank)

    # Initialize RPC.
    rpc.init_rpc(
        name=name,
        rank=rank,
        world_size=world_size
    )

    # Use a distributed autograd context.
    with dist_autograd.context() as context_id:
        # Forward pass (create references on remote nodes).
        rref1 = rpc.remote(dst_name, random_tensor)
        rref2 = rpc.remote(dst_name, random_tensor)
        loss = rref1.to_here() + rref2.to_here()

        # Backward pass (run distributed autograd).
        dist_autograd.backward(context_id, [loss.sum()])

        # Build DistributedOptimizer.
        dist_optim = DistributedOptimizer(
        optim.SGD,
        [rref1, rref2],
        lr=0.05,
        )

        # Run the distributed optimizer step.
        dist_optim.step(context_id)

def run_process(rank, world_size):
    dst_rank = (rank + 1) % world_size
    _run_process(rank, dst_rank, world_size)
    rpc.shutdown()

if __name__ == '__main__':
  # Run world_size workers
  world_size = 2
  mp.spawn(run_process, args=(world_size,), nprocs=world_size)

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

Spec-Zone.ru

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