Хуки ускорителей
Дата создания: 13 нояб. 2025 г. | Дата последнего обновления: 09 дек. 2025 г.
Общие сведения
Хуки ускорителей — это механизм интеграции пользовательских устройств-ускорителей в среду выполнения PyTorch.
Проектирование
В таблицах ниже перечислены хуки, которые поставщикам ускорителей следует реализовать при интеграции новой серверной части устройства. Эти хуки разделены на два уровня приоритета:
- Хуки высокого приоритета: основные API, от которых напрямую зависит среда выполнения PyTorch. Для обеспечения совместимости с основными компонентами и базовой функциональности устройства поставщикам следует реализовать все хуки высокого приоритета.
- Хуки низкого приоритета: API управления устройствами и вспомогательные API, от которых PyTorch напрямую не зависит. Эти хуки повышают удобство использования и обеспечивают поддержку нескольких устройств, но являются необязательными. Поставщики могут реализовать их с учетом конкретных требований и сценариев использования.
Хуки высокого приоритета
Метод хука | Описание | Сценарии использования |
|---|---|---|
| Инициализирует среду выполнения ускорителя и контексты устройств | Настраивает необходимое состояние при первом обращении PyTorch к устройству |
| Проверяет, существует ли для устройства основной контекст | Определяет, выполнялась ли инициализация устройства |
| Возвращает генератор случайных чисел по умолчанию для устройства | Обеспечивает доступ к основному генератору случайных чисел устройства для воспроизводимых случайных операций |
| Создает новый независимый генератор случайных чисел | Создает изолированные экземпляры генераторов для параллельных операций |
| Определяет, какому устройству принадлежит указатель на память | Определяет устройство-ускоритель, связанное с выделенной областью памяти |
| Возвращает распределитель для закрепленной (невыгружаемой) памяти хоста | Выделяет память хоста, которую можно эффективно передавать на ускоритель и обратно |
| Проверяет, указывает ли указатель на закрепленную память | Проверяет типы памяти перед выполнением операций |
Хуки низкого приоритета
Метод хука | Описание | Сценарии использования |
|---|---|---|
| Возвращает значение, указывающее, собрана ли серверная часть ускорителя и включена ли она в расширение | Проверяет, доступна ли библиотека ускорителя во время компиляции |
| Возвращает значение, указывающее, доступно ли оборудование ускорителя во время выполнения | Проверяет, можно ли обнаружить и инициализировать устройства-ускорители |
| Возвращает количество доступных устройств-ускорителей | Перечисляет все доступные устройства-ускорители для выбора устройства |
| Устанавливает активное устройство для текущего потока | Переключает контекст текущего потока на указанное устройство-ускоритель |
| Возвращает индекс текущего активного устройства | Позволяет узнать, какое устройство-ускоритель активно в текущем потоке |
| Атомарно заменяет текущее устройство и возвращает предыдущее | Временно переключает устройство и затем восстанавливает предыдущее |
| Условно заменяет устройство, только если индекс допустим | Безопасно выполняет попытку переключения устройства с проверкой |
Реализация
В качестве примера рассмотрим OpenReg (Open Registration) — интеграционный проект для PyTorch, который восполняет пробел в интеграции серверных частей ускорителей вне основного дерева исходного кода. Он показывает, как поставщики могут регистрировать пользовательские серверные части устройств, не изменяя ядро PyTorch, реализуя интерфейс хуков (см. at::PrivateUse1HooksInterface).
В качестве примера используем getDefaultGenerator:
1 const at::Generator& getDefaultGenerator(DeviceIndex device_index) const override {
2 return getDefaultOpenRegGenerator(device_index);
3 }
В этой реализации:
-
Переопределение базового интерфейса: Метод
getDefaultGeneratorпереопределяет виртуальный метод изat::PrivateUse1HooksInterface. -
Передача управления реализации для конкретного устройства: Вызовите
getDefaultOpenRegGenerator(device_index), который управляет экземпляром генератора для каждого устройства. -
Возврат генератора для конкретного устройства: Возвращаемый
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