进阶 pytorch.org 2026-10-07 23:07:35 · 6 阅读

第10章 利用 PrivateUse1 机制将新硬件后端集成到 PyTorch

Facilitating New Backend Integration by PrivateUse1#Created On: Oct 03, 2023 | Last Updated: May 14, 2026 | Last Verified: Nov 05, 2024

本教程将介绍如何通过 PrivateUse1 机制集成 pytorch/pytorch 仓库之外的新后端所需的必要步骤。假设你已具备 PyTorch 的基础知识。

注意:本教程仅涉及 PrivateUse1 机制中与集成新设备相关的部分,其他内容不予涵盖。同时,教程中提到的所有模块并非必须实现,你可以根据实际需求选择性实现有用的模块。

什么是 PrivateUse1?#

在 PyTorch 2.0 之前,PyTorch 提供了三个预留的 dispatch key(及其对应的 Autograd key)用于原型开发树外后端扩展。这三个 dispatch key 如下:

  • PrivateUse1 / AutogradPrivateUse1
  • PrivateUse2 / AutogradPrivateUse2
  • PrivateUse3 / AutogradPrivateUse3

原型验证通过后,新后端可以申请专用 key,例如 CUDA、XLA、MPS 等。然而,随着 PyTorch 的快速发展和越来越多硬件厂商试图将其后端集成到 PyTorch 中,出现了以下问题:

  • 每个新后端的集成涉及大量文件修改
  • 目前 DispatchKey 数量存在硬性限制(DispatchKeySet 的 64 位限制)

通过 PrivateUse1 key 集成新后端也存在一个问题:无法同时集成多个后端。幸运的是,这些树外后端很少被同时使用。

鉴于上述原因,社区开始推荐通过 PrivateUse1 将新后端集成到 PyTorch 中。然而,旧的 PrivateUse1 机制无法完全胜任新后端集成,因为它在 Storage、AMP、分布式等某些模块中缺乏相关支持。随着 PyTorch 2.1.0 的发布,PrivateUse1 在新后端集成方面进行了一系列优化和增强,现在能够支持快速、高效地集成新设备。

如何通过 PrivateUse1 集成新后端 #

本节将讨论通过 PrivateUse1 将新后端集成到 PyTorch 中的细节,主要包含以下部分:

  • 为新后端注册内核(kernels)。
  • 为新后端注册生成器(generator)。
  • 为新后端注册设备守卫(device guard)。
  • 为新后端元数据注册序列化和反序列化函数。
  • 其他模块。

为新后端注册内核 #

新后端可能具有高性能的算子实现,可以通过 C++ 中描述的 Registering a Dispatched Operator 里的 TORCH_LIBRARY_IMPL API 注册到分发器(dispatcher)中。这涉及以下几种情况:

  • 向分发器注册新后端支持的所有前向算子,同时注册 fallback,以便在新后端不支持某些算子时,这些算子可以回退到 CPU 执行,确保功能可用性。

at::Tensor wrapper_Custom_Tensor_add(const at::Tensor & self, const at::Tensor & other, const at::Scalar & alpha) { // 新后端中 add 内核的实现 ... }

TORCH_LIBRARY_IMPL(aten, PrivateUse1, m) { ... m.impl("add.Tensor", TORCH_FN(wrapper_Custom_Tensor_add)); ... }

void custom_cpu_fallback(const c10::OperatorHandle& op, torch::jit::Stack* stack) { // 添加关于不支持该操作且需回退到 CPU 的新设备的提示 at::native::cpu_fallback(op, stack); }

TORCH_LIBRARY_IMPL(_, PrivateUse1, m) { m.fallback(torch::CppFunction::makeFromBoxedFunction<&custom_cpu_fallback>()); }

  • 如果新后端需要覆盖 PyTorch Autograd 层,通过 AutogradPrivateUse1 将来自 torch::autograd::Function 的内核注册到分发器。分发器和 autograd 系统将自动调用这些算子的前向和反向实现。

class CumtomSeluFunction : public torch::autograd::Function { // 新后端中 selu 内核的实现 }

at::Tensor wrapper_AutogradCumstom__selu(const at::Tensor & self) { return CumtomSeluFunction::apply(self); }

TORCH_LIBRARY_IMPL(aten, AutogradPrivateUse1, m) { ... m.impl("selu", TORCH_FN(wrapper_AutogradCustom__selu)); ... }

  • 通过 AutocastPrivateUse1 将支持自动混合精度(AMP)和回退机制的内核注册到分发器。autocast 系统会在需要时自动调用这些内核。

TORCH_LIBRARY_IMPL(aten, AutocastPrivateUse1, m) { ... KERNEL_PRIVATEUSEONE(<operator>, <policy>) ... }

TORCH_LIBRARY_IMPL(_, AutocastPrivateUse1, m) { m.fallback(torch::CppFunction::makeFallthrough()); }

需要补充的是,如果希望新后端支持 AMP,需要通过 torch._register_device_module("backend_name", BackendModule) 注册新的 BackendModule,并且 BackendModule 需要具有以下 API:

  • get_amp_supported_dtype() -> List[torch.dtype]:获取 AMP 在新后端中支持的 dtypes,可能多支持一个 dtype。
  • is_autocast_enabled() -> bool:检查 AMP 是否在新后端中启用。
  • get_autocast_dtype() -> torch.dtype:获取 AMP 在新后端中支持的 dtype,该值由 set_autocast_dtype 或默认 dtype 设置,默认 dtype 为 torch.float16。
  • set_autocast_enabled(bool) -> None:在新后端中启用或禁用 AMP。
  • set_autocast_dtype(dtype) -> None:设置 AMP 在新后端中支持的 dtype,该 dtype 必须包含在 get_amp_supported_dtype 返回的 dtypes 中。

为新后端注册生成器 #

需要支持对应于新设备的生成器。目前,PrivateUse1 可以动态注册自定义生成器,主要步骤如下。

  1. 继承 GeneratorImpl 类实现新后端对应的生成器类,并实现各种通用方法。
  2. 定义一个新的后端构建器,带有单个参数:device index。
  3. 调用 REGISTER_GENERATOR_PRIVATEUSE1 宏完成动态注册。

struct CustomGeneratorImpl : public c10::GeneratorImpl { // 新后端中生成器的实现 }

at::Generator make_custom_generator(c10::DeviceIndex device_index) { return at::make_generator(device_index); }

REGISTER_GENERATOR_PRIVATEUSE1(make_cumstom_generator)

为新后端注册设备守卫 #

PyTorch 通过 DeviceGuard 提供与设备、流和事件切换相关的功能。该功能也适用于 PrivateUse1 Key。

  1. 继承 DeviceGuardImplInterface 类实现新后端对应的各种通用方法。
  2. 调用 C10_REGISTER_GUARD_IMPL 宏完成动态注册。

struct CustomGuardImpl final : public c10::impl::DeviceGuardImplInterface { // 新后端中守卫的实现 }

C10_REGISTER_GUARD_IMPL(PrivateUse1, CustomGuardImpl);

为新后端元数据注册序列化和反序列化函数 #

PyTorch 目前能够动态注册序列化/反序列化函数,以支持 TensorImpl.ExtraMeta 类中名为 backend_meta_ 的新后端附加元数据的序列化和反序列化。你可以参考以下步骤:

  1. 继承 BackendMeta 类实现对应于新后端的 CustomBackendMetadata,可以在该类中自定义新后端的各个字段。
  2. 实现新后端的序列化和反序列化函数,函数签名为 void(const at::Tensor&, std::unordered_map&)。
  3. 调用 TensorBackendMetaRegistry 宏完成动态注册。

struct CustomBackendMetadata : public c10::BackendMeta { // 新后端中后端元数据的实现 }

void for_serialization(const at::Tensor& t, std::unordered_map& m) { // 序列化的实现 }

void for_deserialization(const at::Tensor& t, std::unordered_map& m) { // 反序列化的实现 }

TensorBackendMetaRegistry(c10::DeviceType::PrivateUse1, &for_serialization, &for_deserialization);

其他模块 #

除上述部分外,还有一些其他模块可以通过 PrivateUse1 进行扩展,例如分布式集合通信、基准计时器等,这些将在未来添加。Ascend NPU 是一个关于 PrivateUse1 集成的例子。

如何利用 PrivateUse1 改善用户体验 #

通过 PrivateUse1 集成新设备的首要目标是满足基本功能需求,下一步则是提高易用性,主要涉及以下方面。

  • 向 PyTorch 注册新后端模块。
  • 将 PrivateUse1 重命名为新后端的自定义名称。
  • 生成与新后端相关的方法和属性。

向 PyTorch 注册新后端模块 #

PyTorch 中一些与 CUDA 相关的接口可以通过以下形式调用:torch.cuda.xxx。因此,为了符合用户习惯,通过 PrivateUse1 机制实现的新后端也应提供类似的接口。例如,使用 Ascend NPU:

torch._register_device_module('npu', torch_npu.npu)

完成上述操作后,用户可以通过 torch.npu.xxx 调用 Ascend NPU 的专用 API。

将 PrivateUse1 重命名为新后端的自定义名称 #

PrivateUse1 Key 是集成到 PyTorch 中的新后端的内部机制。对于用户来说,与 PrivateUse1 相比,与新后端强相关的自定义名称应更友好。以 Ascend NPU 为例,第一种用法对用户更友好:

torch.rand((2,2),device='npu:0')

torch.rand((2,2),device='privateuse1:0')

现在,PyTorch 为自命名的 PrivateUse1 后端提供了新的 C++/Python API,使用非常简单。

PYTHON torch.rename_privateuse1_backend("npu")

C++ c10::register_privateuse1_backend("npu")

生成与新后端相关的方法和属性 #

将 PrivateUse1 重命名为自定义名称后,自动为 Tensor、nn、Storage 模块生成与新后端名称相关的属性和方法。以下是 Ascend NPU 的示例:

torch.rename_privateuse1_backend("npu") unsupported_dtype = [torch.quint8] torch.utils.generate_methods_for_privateuse1_backend(for_tensor=True, for_module=True, for_storage=True, unsupported_dtype=unsupported_dtype)

然后,你可以使用以下方法和属性:

torch.Tensor.npu() torch.Tensor.is_npu torch.Storage.npu() torch.Storage.is_npu ...

未来工作 #

PrivateUse1 机制的改进仍在进行中,新模块的 PrivateUse1 集成方法将陆续添加。以下是我们正在积极工作的几个项目:

  • 添加分布式集合通信的集成方法。
  • 添加基准计时器的集成方法。

总结 #

本教程带你通过 PrivateUse1 将新后端集成到 PyTorch 中的过程,包括但不限于算子注册、生成器注册、设备守卫注册等。同时,还介绍了一些改善用户体验的方法。

评论 (0)