Spec-Zone.ru › PyTorch 2.14

Хуки ускорителей

Дата создания: 13 нояб. 2025 г. | Дата последнего обновления: 09 дек. 2025 г.

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

Хуки ускорителей — это механизм интеграции пользовательских устройств-ускорителей в среду выполнения PyTorch.

Проектирование

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

  • Хуки высокого приоритета: основные API, от которых напрямую зависит среда выполнения PyTorch. Для обеспечения совместимости с основными компонентами и базовой функциональности устройства поставщикам следует реализовать все хуки высокого приоритета.
  • Хуки низкого приоритета: API управления устройствами и вспомогательные API, от которых PyTorch напрямую не зависит. Эти хуки повышают удобство использования и обеспечивают поддержку нескольких устройств, но являются необязательными. Поставщики могут реализовать их с учетом конкретных требований и сценариев использования.

Хуки высокого приоритета

Метод хука

Описание

Сценарии использования

init()

Инициализирует среду выполнения ускорителя и контексты устройств

Настраивает необходимое состояние при первом обращении PyTorch к устройству

hasPrimaryContext(DeviceIndex)

Проверяет, существует ли для устройства основной контекст

Определяет, выполнялась ли инициализация устройства

getDefaultGenerator(DeviceIndex)

Возвращает генератор случайных чисел по умолчанию для устройства

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

getNewGenerator(DeviceIndex)

Создает новый независимый генератор случайных чисел

Создает изолированные экземпляры генераторов для параллельных операций

getDeviceFromPtr(void*)

Определяет, какому устройству принадлежит указатель на память

Определяет устройство-ускоритель, связанное с выделенной областью памяти

getPinnedMemoryAllocator()

Возвращает распределитель для закрепленной (невыгружаемой) памяти хоста

Выделяет память хоста, которую можно эффективно передавать на ускоритель и обратно

isPinnedPtr(void*)

Проверяет, указывает ли указатель на закрепленную память

Проверяет типы памяти перед выполнением операций

Хуки низкого приоритета

Метод хука

Описание

Сценарии использования

isBuilt()

Возвращает значение, указывающее, собрана ли серверная часть ускорителя и включена ли она в расширение

Проверяет, доступна ли библиотека ускорителя во время компиляции

isAvailable()

Возвращает значение, указывающее, доступно ли оборудование ускорителя во время выполнения

Проверяет, можно ли обнаружить и инициализировать устройства-ускорители

deviceCount()

Возвращает количество доступных устройств-ускорителей

Перечисляет все доступные устройства-ускорители для выбора устройства

setCurrentDevice(DeviceIndex)

Устанавливает активное устройство для текущего потока

Переключает контекст текущего потока на указанное устройство-ускоритель

getCurrentDevice()

Возвращает индекс текущего активного устройства

Позволяет узнать, какое устройство-ускоритель активно в текущем потоке

exchangeDevice(DeviceIndex)

Атомарно заменяет текущее устройство и возвращает предыдущее

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

maybeExchangeDevice(DeviceIndex)

Условно заменяет устройство, только если индекс допустим

Безопасно выполняет попытку переключения устройства с проверкой

Реализация

В качестве примера рассмотрим OpenReg (Open Registration) — интеграционный проект для PyTorch, который восполняет пробел в интеграции серверных частей ускорителей вне основного дерева исходного кода. Он показывает, как поставщики могут регистрировать пользовательские серверные части устройств, не изменяя ядро PyTorch, реализуя интерфейс хуков (см. at::PrivateUse1HooksInterface).

В качестве примера используем getDefaultGenerator:

1  const at::Generator& getDefaultGenerator(DeviceIndex device_index) const override {
2    return getDefaultOpenRegGenerator(device_index);
3  }

В этой реализации:

  1. Переопределение базового интерфейса: Метод getDefaultGenerator переопределяет виртуальный метод из at::PrivateUse1HooksInterface.
  2. Передача управления реализации для конкретного устройства: Вызовите getDefaultOpenRegGenerator(device_index), который управляет экземпляром генератора для каждого устройства.
  3. Возврат генератора для конкретного устройства: Возвращаемый at::Generator оборачивает OpenRegGeneratorImpl, реализующий генерацию случайных чисел для конкретного устройства.

Этот подход применим ко всем хукам: переопределите метод интерфейса, проверьте входные данные, передайте управление API для конкретного устройства и верните результаты в ожидаемом PyTorch формате.

Пример интеграции

В следующих разделах показано, как PyTorch взаимодействует с хуками ускорителей при получении генератора случайных чисел по умолчанию. В примере прослеживается полный путь от пользовательского кода на Python до реализации для конкретного устройства.

Уровень 1: пользовательский код

Пользовательский код задает детерминированное начальное значение, вызывая manual_seed:

import torch
torch.openreg.manual_seed(42)

Уровень 2: API расширения Python

Уровень API Python управляет выбором устройства и вызывает расширение C++ (определенное в torch_openreg/openreg/random.py):

1def manual_seed(seed: int) -> None:
2    seed = int(seed)
3
4    idx = current_device()
5    default_generator = torch_openreg._C._get_default_generator(idx)
6    default_generator.manual_seed(seed)
7
8

Функция manual_seed получает индекс текущего устройства, вызывает torch_openreg._C._get_default_generator(idx) для получения генератора для конкретного устройства и задает начальное значение.

Уровень 3: мост Python/C++

Расширение C++ предоставляет Python функцию _getDefaultGenerator, которая служит мостом к ядру PyTorch:

 1static PyObject* _getDefaultGenerator(PyObject* self, PyObject* arg) {
 2  HANDLE_TH_ERRORS
 3  TORCH_CHECK(
 4      THPUtils_checkLong(arg),
 5      "_get_default_generator expects an int, but got ",
 6      THPUtils_typename(arg));
 7  auto idx = static_cast<int>(THPUtils_unpackLong(arg));
 8
 9  torch::utils::register_fork_handler_for_device_init(at::kPrivateUse1);
10  return THPGenerator_initDefaultGenerator(
11      at::globalContext().defaultGenerator(
12          c10::Device(c10::DeviceType::PrivateUse1, idx)));
13
14  END_HANDLE_TH_ERRORS
15}
1static PyMethodDef methods[] = {
2    {"_init", _initExtension, METH_NOARGS, nullptr},
3    {"_isInBadFork", _isInBadFork, METH_NOARGS, nullptr},
4    {"_get_default_generator", _getDefaultGenerator, METH_O, nullptr},
5    {"_get_device", _getDevice, METH_NOARGS, nullptr},
6    {"_set_device", _setDevice, METH_O, nullptr},
7    {"_exchangeDevice", _exchangeDevice, METH_O, nullptr},
8    {"_get_device_count", _getDeviceCount, METH_NOARGS, nullptr},
9    {nullptr, nullptr, 0, nullptr}};

Эта функция извлекает из Python индекс устройства, создает объект устройства PrivateUse1 и вызывает at::globalContext().defaultGenerator(). Затем контекст PyTorch направляет вызов зарегистрированным хукам.

Уровень 4: контекст ядра PyTorch

Класс Context PyTorch направляет вызовы соответствующим хукам ускорителя (aten/src/ATen/Context.h):

 1TORCH_API std::string precision2str(Float32Precision prec);
 2TORCH_API CuDNNDepthwiseKernel str2cudnn_depthwise(const std::string& name);
 3TORCH_API std::string cudnn_depthwise2str(CuDNNDepthwiseKernel k);
 4
 5class TORCH_API Context {
 6 public:
 7  Context();
 8
 9  const Generator& defaultGenerator(Device device) {
10    c10::DeviceType device_type = device.type();
11    lazyInitDevice(device_type);
12
13    if (device_type == at::kCPU) {
14      return at::detail::getDefaultCPUGenerator();
15    } else {
16      return getAcceleratorHooksInterface(device_type)
17          .getDefaultGenerator(device.index());
18    }
19  }
20
21  const AcceleratorHooksInterface& getAcceleratorHooksInterface(
22      std::optional<c10::DeviceType> opt_device_type = std::nullopt) {
23    if (!opt_device_type.has_value()) {
24      opt_device_type = at::getAccelerator(true);
25    }
26    if (opt_device_type == at::kCUDA) {
27      return at::detail::getCUDAHooks();
28    } else if (opt_device_type == at::kXPU) {
29      return at::detail::getXPUHooks();
30    } else if (opt_device_type == at::kMPS) {
31      return at::detail::getMPSHooks();
32    } else if (opt_device_type == at::kPrivateUse1) {
33      return at::detail::getPrivateUse1Hooks();
34    } else if (opt_device_type == at::kMTIA) {
35      return at::detail::getMTIAHooks();
36    } else if (opt_device_type == at::kHIP) {
37      return at::detail::getHIPHooks();
38    } else if (opt_device_type == at::kHPU) {
39      return at::detail::getHPUHooks();
40    } else if (opt_device_type == at::kXLA) {
41      return at::detail::getXLAHooks();
42    } else {
43      TORCH_CHECK(
44          false,

Такая многоуровневая архитектура позволяет PyTorch не зависеть от конкретного устройства, передавая операции, зависящие от оборудования, реализациям ускорителей. Хуки регистрируются один раз при загрузке модуля:

1namespace c10::openreg {
2
3static bool register_hook_flag [[maybe_unused]] = []() {
4  at::RegisterPrivateUse1HooksInterface(new OpenRegHooksInterface());
5
6  return true;
7}();
8
9} // namespace c10::openreg

Уровень 5: хуки ускорителей

Интерфейс хуков предоставляет абстракцию, с помощью которой PyTorch передает управление реализациям для конкретных устройств:

1  const at::Generator& getDefaultGenerator(DeviceIndex device_index) const override {
2    return getDefaultOpenRegGenerator(device_index);
3  }

Метод хука getDefaultGenerator переопределяет базовый интерфейс и передает управление getDefaultOpenRegGenerator, который управляет фактическими экземплярами генераторов.

Уровень 6: реализация для конкретного устройства

Реализация для конкретного устройства управляет экземплярами генераторов для каждого устройства:

 1const at::Generator& getDefaultOpenRegGenerator(c10::DeviceIndex device_index) {
 2  static bool flag [[maybe_unused]] = []() {
 3    auto deivce_nums = device_count();
 4    default_generators.resize(deivce_nums);
 5    for (auto i = 0; i < deivce_nums; i++) {
 6      default_generators[i] = at::make_generator<OpenRegGeneratorImpl>(i);
 7      default_generators[i].seed();
 8    }
 9    return true;
10  }();
11
12  c10::DeviceIndex idx = device_index;
13  if (idx == -1) {
14    idx = current_device();
15  } else {
16    TORCH_CHECK(idx >= 0 && idx < device_count());
17  }
18  return default_generators[idx];
19}

Эта функция поддерживает статический вектор генераторов (по одному на устройство), инициализирует их при первом обращении, проверяет индекс устройства и возвращает соответствующий экземпляр генератора.

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

Spec-Zone.ru

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