评价此页

自定义 Python 算子#

创建日期:2026年6月16日 | 最后更新:2026年6月16日 | 最后验证:2024年11月5日

您将学到什么
  • 何时创建自定义 Python 算子

  • 如何在函数式算子契约与可变算子契约之间做出选择

  • 为什么需要架构(schema)和变异/别名契约

  • 伪内核(fake kernels)、自动微分及其他注册信息的适用场景

PyTorch 提供了大量的张量运算库,例如 torch.addtorch.sum。然而,您可能希望在 PyTorch 中使用新的自定义算子,这些算子可能由第三方库编写。本指南展示了如何包装 Python 函数,使其行为表现得如同 PyTorch 原生算子一样。

您可能希望在 PyTorch 中创建自定义算子的原因包括:

  • torch.compile 和/或 torch.export 中将任意 Python 函数视为不透明的可调用对象(opaque callable);

  • 为任意 Python 函数添加训练支持。

请注意,如果您的操作可以表示为现有 PyTorch 算子的组合,则通常无需使用自定义算子 API。torch.compile、训练支持以及其他 PyTorch 子系统通常可以直接工作。

每个自定义算子都需要:

  • 稳定的架构和变异/别名契约;

  • 通过 torch.library.opcheck 进行验证;

  • 如果算子返回张量且必须与 torch.compiletorch.export 一起使用,则需要伪内核。

选择路径

  • 函数式自定义算子:算子返回全新的张量且不修改任何输入。

  • 可变自定义算子:算子修改输入或写入输出缓冲区。从 PyTorch 2.13 开始,这包括 PyTorch 风格的就地(in-place)和 out= 自定义算子。

  • 可选注册:在基础算子通过 opcheck 后,添加自动微分、torch.vmap、张量子类行为或其他子系统支持。

选择您的路径
  • 任何自定义算子:阅读 架构和变异/别名契约 以及 验证。您需要稳定的架构、代表性的示例以及 opcheck

  • 返回新张量且不修改输入的代码:阅读 函数式自定义算子。您需要 custom_op(..., mutates_args=())、用于 torch.compile 的伪内核,以及 opcheck

  • 写入现有内存的内核:阅读 可变自定义算子。您需要准确的 mutates_args 和一个明确的变异模式。

  • 就地(in-place)、``out=`` 或“也许输出(maybe-out)”行为:阅读 可变自定义算子架构契约。从 PyTorch 2.13 开始,提供已标记的就地和 out= 自定义算子;请将“也许输出”行为拆分为单独的算子。

  • 训练支持、``vmap`` 或张量子类行为:阅读 添加注册。先从经过验证的基础算子开始,然后为该子系统添加注册。

对于非 Python 环境或 AOTInductor,请改用 C++ 定义算子和后端内核。请参阅 C++ 自定义算子教程

开始之前#

内核(kernel)是具体实现。算子(operator)是面向 PyTorch 的契约:名称、输入、输出、变异行为和子系统注册。

自定义算子为 PyTorch 提供了一个明确的边界。当追踪(tracing)实现过程不可行或不理想时,请使用它。

必要条件:架构和变异/别名契约#

在编写注册信息之前,请先确定架构和变异/别名契约。PyTorch 使用架构和注册信息来推断变异/别名关系;它不会从 Python 函数体中推断契约。

当两个张量共享相同的底层存储时,它们就会产生别名。例如,y = x.view(-1) 创建了一个 视图(view) y,它与 x 产生别名,因此写入 y 可能会改变 x

  • 架构必须是稳定的:变异和别名行为必须正确且一致。这意味着算子不得返回有时会与输入产生别名的输出。此外,算子不得修改未标记为正在被修改的输入。

  • 函数式自定义算子必须返回全新的张量。不要返回输入张量、输入的视图,或互为别名的两个输出。

  • 可变自定义算子必须在 mutates_args 中列出每个被修改的参数。

  • 伪内核必须返回与真实内核具有相同元数据的张量:形状、数据类型(dtype)、设备、布局、步长(strides)以及适用的存储偏移量。empty_like(x) 仅在真实输出与 x 具有相同元数据时才正确。函数式自定义算子页面展示了一个关于此元数据不匹配的可执行示例。

  • 伪内核可以检查元数据,但绝不能读取张量数据。

  • 避免使用“也许输出(maybe-out)”算子。有时分配新张量而有时写入输出缓冲区的算子,在不同的调用中具有不同的别名契约。

将“也许输出”行为拆分为两个算子:一个负责分配的函数式算子,和一个写入输出缓冲区的可变算子。

必要条件:使用 opcheck 验证#

torch.library.opcheck 用于验证注册契约:架构、伪内核、自动微分注册以及在编译 API 下的行为。

在代表性输入上运行 opcheck

  • 每个支持的设备;

  • 重要的数据类型(dtypes);

  • 边缘形状,例如空张量;

  • 重要的内存格式或非连续步长;

  • 如果算子支持训练,则包括 requires_grad=True 的输入。

opcheck 并非数值正确性测试。请使用 torch.testing.assert_close 或常规单元测试进行前向正确性验证,并使用 torch.autograd.gradcheck 验证梯度公式。

后续步骤#

先阅读一页基础契约页面,仅在需要时添加注册: