PyTorch 加速器算子注册实战:基于 PrivateUse1 的算子适配、Fallback 与 STUB 机制全解析
2026/9/8 22:23:14 网站建设 项目流程

PyTorch 加速器算子注册实战:基于 PrivateUse1 的算子适配、Fallback 与 STUB 机制全解析

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

本篇围绕 PyTorch 加速器集成文档中的"Operator Registration(算子注册)"章节展开,系统讲解新加速器(以PrivateUse1Dispatch Key 为例)如何注册内置算子、启用 CPU Fallback、通过STUB二次调度以及定义自定义算子,并结合仓库中torch_openreg参考实现与aten源码,给出每一步可直接复用的代码与底层机制。读完本篇,读者能够独立完成一个新后端从最小算子集到自定义算子的完整注册链路,并熟练使用torch._C._dispatch_*系列命令与环境变量调试调度过程。

为什么算子注册是加速器集成的核心

对新加速器而言,集成 PyTorch 最基础也最核心的工作就是支持高性能算子。PyTorch 为此提供了PythonC++双栈的算子开发注册方法。理解这一切的前提是Dispatch Key机制:

Dispatch Key用于在 PyTorch 中唯一标识一个加速器,例如CPUCUDAMPSPrivateUse1。理论上,所有后续新增的加速器都共享PrivateUse1这一 Dispatch Key,借助它内置的完整脚手架能力来完成新加速器的集成。

从源码结构看,PrivateUse1的算子分发基础设施位于 DispatchStub.h 中。该头文件为每个平台定义了独立的注册宏,其中与本文主题直接相关的是:

#define REGISTER_PRIVATEUSE1_DISPATCH(name, fn) \ static RegisterPRIVATEUSE1Dispatch<struct name##_DECLARE_DISPATCH_type> name ## __register(name, fn);

(参见 DispatchStub.h)

此外,aten目录下还存在DECLARE_DISPATCH宏,它显式声明了可被二次调度的STUB(本文"STUB"一节会详细展开)。整条加速器集成路径在 Accelerator Integration 文档首页 中有总体介绍,算子章节(operators)是其中"Runtime / Operators / Python Frontend / High-level Modules"四大轴线中的第二条主线。

Operator Set:新后端必须优先实现的算子集

PyTorch 目前有超过 3500 个内置算子(含各种变体)。要在短期内适配全部算子不现实,因此新后端开发的第一步应当聚焦于必需算子集:其余算子先用社区提供的 Fallback 机制兜底保证功能正确,再逐步补全以提升性能。

必需的算子集如下(均为PrivateUse1),主要由工厂函数依赖的底层算子和 fallback 算子构成:

算子名称Dispatch Key说明
empty.memory_formatPrivateUse1创建指定 shape 与内存布局的未初始化 Tensor(stride 自动计算)
empty_stridedPrivateUse1创建指定 shape 与 stride 的未初始化 Tensor(自由度更大)
as_stridedPrivateUse1以新 shape/stride/offset 创建输入 Tensor 的共享视图(不分配新内存)
viewPrivateUse1创建新 shape 的共享视图,但要求原 Tensor 内存连续
_reshape_aliasPrivateUse1无安全检查地创建共享视图(reshape 的内部版本)
resize_PrivateUse1原地修改 Tensor shape,容量不足时重新分配内存
_copy_fromPrivateUse1Tensor.copy_的底层核心函数,负责实际的跨设备数据拷贝
_copy_from_and_resizePrivateUse1组合resize__copy_from,先 resize 再拷贝
_local_scalar_densePrivateUse1.item()的底层实现,从 Tensor 提取 CPU 标量值
set_.source_TensorPrivateUse1使用指定 Tensor 设置当前 Tensor
set_.source_StoragePrivateUse1使用指定 Storage 设置当前 Tensor
set_.source_Storage_storage_offsetPrivateUse1使用指定 Storage 及其 storage offset 设置当前 Tensor
fallbackPrivateUse1Fallback 到 CPU 执行

仓库中的参考实现 Minimal.cpp 恰好实现了上表中的绝大部分算子,可作为新后端起步的样板代码。

Step 1:实现并注册内置算子(以 empty.memory_format 为例)

上述算子有一个共同特点:它们都是 PyTorch 内置算子,拥有已定义的namespaceSchema,且内置加速器(CPUCUDA等)已有实现。我们只需要为新加速器补齐实现。

1.1 在 native_functions.yaml 中查询 Schema

empty.memory_format为例,第一步是查询该算子的 schema 信息(包含完整签名)。当前仓库中其定义位于 native_functions.yaml:

- func: empty.memory_format(SymInt[] size, *, ScalarType? dtype=None, Layout? layout=None, Device? device=None, bool? pin_memory=None, MemoryFormat? memory_format=None) -> Tensor dispatch: CPU: empty_cpu CUDA: empty_cuda ...

从该文件可以看到CPU对应empty_cpuCUDA对应empty_cuda的分发关系。新后端需要做的就是按照同样签名,为新加速器设备写一份实现。

1.2 编写算子实现

参考实现位于 Minimal.cpp:

at::Tensor empty_memory_format( c10::IntArrayRef size, std::optional<c10::ScalarType> dtype_opt, std::optional<c10::Layout> layout_opt, std::optional<c10::Device> device_opt, std::optional<bool> pin_memory_opt, std::optional<c10::MemoryFormat> memory_format_opt) { const auto device = c10::device_or_default(device_opt); const auto dtype = c10::dtype_or_default(dtype_opt); TORCH_CHECK(device.is_privateuseone()); TORCH_CHECK( c10::layout_or_default(layout_opt) == c10::Layout::Strided, "Non strided layout not supported"); TORCH_CHECK( !c10::pinned_memory_or_default(pin_memory_opt), "Pin memory can only be on CPU"); const c10::DeviceGuard device_guard(device); constexpr c10::DispatchKeySet pu1_dks(c10::DispatchKey::PrivateUse1); auto allocator = at::GetAllocator(at::kPrivateUse1); return at::detail::empty_generic( size, allocator, pu1_dks, dtype, memory_format_opt); }

几个关键细节:

  • 通过c10::device_or_default/c10::dtype_or_default填充默认值,这是 PyTorch 算子实现的通用模式;
  • TORCH_CHECK明确拒绝非PrivateUse1设备、非Strided布局和 pinned memory 等不支持的组合,让错误尽早暴露;
  • DeviceGuard保证后续内存分配发生在正确设备上;
  • 通过at::GetAllocator(at::kPrivateUse1)获取设备分配器后调用at::detail::empty_generic完成实际分配——这正是注册设备分配器(设备集成文档中 Device 章节)与算子注册的衔接点。

同文件中empty_stridedas_stridedresize__reshape_alias_copy_from_copy_from_and_resize_local_scalar_denseset_.source_*view的实现思路一致(参见 Minimal.cpp),其中_copy_from还展示了跨设备拷贝如何通过at::from_blob构造 CPU 视角的张量后调用at::native::copy_来完成。

1.3 编写 Wrapper 并通过 TORCH_LIBRARY_IMPL 注册

实现完成后,需要一个与 schema 签名严格一致的 wrapper 函数,再通过TORCH_LIBRARY_IMPLaten::empty.memory_format注册到PrivateUse1。参考代码位于 OpenRegMinimal.cpp 与 OpenRegMinimal.cpp:

// LITERALINCLUDE START: EMPTY.MEMORY_FORMAT WRAPPER at::Tensor wrapper_empty_memory_format( c10::IntArrayRef size, std::optional<c10::ScalarType> dtype_opt, std::optional<c10::Layout> layout_opt, std::optional<c10::Device> device_opt, std::optional<bool> pin_memory_opt, std::optional<c10::MemoryFormat> memory_format_opt) { return at::native::openreg::empty_memory_format( size, dtype_opt, layout_opt, device_opt, pin_memory_opt, memory_format_opt); } // LITERALINCLUDE END: EMPTY.MEMORY_FORMAT WRAPPER
// LITERALINCLUDE START: TORCH_LIBRARY_IMPL DEFAULT TORCH_LIBRARY_IMPL(aten, PrivateUse1, m) { m.impl("empty.memory_format", wrapper_empty_memory_format); m.impl("empty_strided", wrapper_empty_strided); m.impl("as_strided", wrapper_as_strided); m.impl("resize_", wrapper_resize_); m.impl("_reshape_alias", wrapper__reshape_alias); m.impl("_copy_from", wrapper__copy_from); m.impl("_copy_from_and_resize", wrapper__copy_from_and_resize); m.impl("_local_scalar_dense", wrapper__local_scalar_densor); m.impl( "_has_compatible_shallow_copy_type", wrapper_has_compatible_shallow_copy_type); m.impl("set_.source_Tensor", wrapper_set_source_Tensor_); m.impl("set_.source_Storage", wrapper_set_source_Storage_); m.impl( "set_.source_Storage_storage_offset", wrapper_set_source_Storage_storage_offsetset_); m.impl("view", wrapper_view); } // LITERALINCLUDE END: TORCH_LIBRARY_IMPL DEFAULT

这里TORCH_LIBRARY_IMPL(aten, PrivateUse1, m)的两个参数分别指定命名空间与 Dispatch Key,m.impl("<算子名>", <wrapper>)逐条完成绑定。注意实际注册中还包含了上表之外但实践中同样重要的_has_compatible_shallow_copy_type(返回true表示允许跨设备的浅拷贝类型兼容),这体现了"从最小算子集起步,按需扩展"的策略。

Step 2:注册全局 Fallback

完成 Step 1 后,必需算子集(除fallback外)已全部就绪。要支持数学运算、卷积等尚未实现的算子,需要注册fallback语义——这是 PyTorch 框架内置的能力,可把新加速器不支持的操作自动回退到 CPU 执行。对新开发中的后端而言,这是"牺牲性能、保证功能"的高效手段。

2.1 编写 fallback 实现

参考实现位于 Minimal.cpp,其核心是包装 PyTorch 提供的at::native::cpu_fallback(声明于 CPUFallback.h):

// LITERALINCLUDE START: FALLBACK IMPL void cpu_fallback(const c10::OperatorHandle& op, torch::jit::Stack* stack) { static const std::unordered_set<c10::OperatorName> cpu_fallback_blocklist = { c10::OperatorName("aten::abs", ""), c10::OperatorName("aten::abs", "out"), }; const auto& op_name = op.schema().operator_name(); if (cpu_fallback_blocklist.count(op_name)) { TORCH_CHECK( false, "Operator '", op_name, "' is not implemented for device openreg."); } else { at::native::cpu_fallback(op, stack); } } // LITERALINCLUDE END: FALLBACK IMPL

op是被调用的算子句柄(可从中取到schema与算子名),stack是 TorchScript 执行栈,at::native::cpu_fallback会负责把参数搬到 CPU、执行对应 CPU kernel、再把结果搬回原设备。

2.2 Wrapper 与全局注册

随后编写 wrapper 并通过m.fallback注册为所有算子的默认实现(参见 OpenRegMinimal.cpp):

// LITERALINCLUDE START: FALLBACK WRAPPER void wrapper_cpu_fallback( const c10::OperatorHandle& op, torch::jit::Stack* stack) { at::native::openreg::cpu_fallback(op, stack); } // LITERALINCLUDE END: FALLBACK WRAPPER
// LITERALINCLUDE START: FALLBACK GLOBAL TORCH_LIBRARY_IMPL(_, PrivateUse1, m) { m.fallback( torch::CppFunction::makeFromBoxedFunction<&wrapper_cpu_fallback>()); } // LITERALINCLUDE END: FALLBACK GLOBAL

注册完成后,新后端不支持的算子会自动回退到 CPU 执行,结果再传回新后端。注意两点:

  • 命名空间参数写作_表示"任意命名空间",即对所有算子生效;
  • torch::CppFunction::makeFromBoxedFunction将自由函数包装成可被 Dispatcher 调用的 boxed function,这是注册 fallback 的标准写法。

Advanced 一:Selective Fallback(选择性回退)

有时我们只希望部分算子启用 fallback,其余算子保持 PyTorch 默认行为(设备没有对应实现就报错)——这在"绝大多数算子已适配、个别算子想先兜底"的场景下非常合理。

选择性 fallback 与全局 fallback 的唯一区别在于注册方式:

  • m.impl("<算子名>", ...):为特定算子注册实现;
  • m.fallback(...):为所有算子注册默认实现。

对单个算子启用 fallback 的注册方式(参见 OpenRegMinimal.cpp):

// LITERALINCLUDE START: FALLBACK SINGLE TORCH_LIBRARY_IMPL(aten, PrivateUse1, m) { m.impl( "sub.Tensor", torch::CppFunction::makeFromBoxedFunction<&wrapper_cpu_fallback>()); } // LITERALINCLUDE END: FALLBACK SINGLE

此外,全局 fallback + 黑名单也是常见组合:当只有少数算子不支持 fallback 时,在 fallback 实现内部维护一个黑名单集合,命中黑名单的算子直接抛错,其余走at::native::cpu_fallback。前面 Step 2 中cpu_fallback实现里的cpu_fallback_blocklist(示例黑名单为aten::absaten::abs.out,见 Minimal.cpp)正是这一模式的落地写法——对黑名单中的算子执行TORCH_CHECK(false, ...)报错,其余算子交给at::native::cpu_fallback(op, stack)处理。

Advanced 二:PyTorch STUB 注册路径

除了TORCH_LIBRARY_IMPL,PyTorch 还提供另一种内置算子注册方式:STUB。它本质上仍基于 Step 1 的路径,但增加了二次调度能力(例如根据 CPU 特性进一步分发到 AVX512/SVE 等架构 kernel)。

需要注意的限制:

  • STUB方式目前仅支持有限集合的算子,PyTorch 并未明确列出可通过STUB注册的算子清单;
  • 对新加速器设备而言,STUB的优势是用较小的性能开销显著降低开发成本——算子可以直接基于TensorIteratorBase等高层抽象编写。

3.1 查找支持 STUB 的算子

DECLARE_DISPATCH是显式声明STUB的宏,当前分布在aten目录中。文档给出的查询命令如下:

pushd ${TORCH_ROOT} find aten -type f -a -name "*.h" | xargs -I {} grep -wl "^DECLARE_DISPATCH" {} popd

它会列出所有声明了 STUB 的声明文件,例如:

... aten/src/ATen/native/Activation.h aten/src/ATen/native/FusedSGD.h aten/src/ATen/native/nested/NestedTensorBinaryOps.h aten/src/ATen/native/TensorCompare.h aten/src/ATen/native/Sorting.h ...

abs_stub为例,其声明位于 UnaryOps.h:

using unary_fn = void(*)(TensorIteratorBase&); DECLARE_DISPATCH(unary_fn, abs_stub)

从签名可以看出 STUB 的输入是TensorIteratorBase——这是 PyTorch 提供的强大辅助类,封装了全部输入/输出算子以及其他辅助方法。

3.2 基于 STUB 实现算子

参考实现分两部分。第一部分是基于TensorIteratorBase的 kernel(参见 Extra.cpp):

// LITERALINCLUDE START: STUB ABS void abs_kernel(at::TensorIteratorBase& iter) { TORCH_CHECK(iter.ntensors() == 2, "Abs kernel expects 2 tensors"); TORCH_CHECK( iter.common_dtype() == at::ScalarType::Float, "Abs kernel only supports float type"); auto& output_tensor = iter.tensor(0); auto& input_tensor = iter.tensor(1); TORCH_CHECK( input_tensor.sizes() == output_tensor.sizes(), "Input and output tensor sizes must match."); auto abs_loop = [](float* out_ptr, const float* in_ptr, int64_t n) { for (int64_t i = 0; i < n; ++i) { out_ptr[i] = std::abs(in_ptr[i]); } }; MemoryGuard guard(input_tensor, output_tensor); if (iter.is_contiguous()) { abs_loop( static_cast<float*>(iter.data_ptr(0)), static_cast<float*>(iter.data_ptr(1)), iter.numel()); } else { TORCH_CHECK( input_tensor.is_contiguous(), "Input tensor must be contiguous.") auto output = at::empty( input_tensor.sizes(), input_tensor.options().memory_format( input_tensor.suggest_memory_format())); MemoryGuard guard(output); abs_loop( static_cast<float*>(output.data_ptr()), static_cast<float*>(iter.data_ptr(1)), iter.numel()); output_tensor.copy_(output); } } // LITERALINCLUDE END: STUB ABS

该示例 kernel 只支持Float类型且输入需为连续布局(源码注释中也明确标注了这一限制)。开发完 kernel 后,第二部分是通过REGISTER_PRIVATEUSE1_DISPATCHabs_stub注册到PrivateUse1(参见 OpenRegExtra.cpp):

// LITERALINCLUDE START: STUB DEFAULT REGISTER_PRIVATEUSE1_DISPATCH(abs_stub, &wrapper_abs_stub); REGISTER_PRIVATEUSE1_DISPATCH( quantize_tensor_per_tensor_affine_stub, &wrapper_quantize_tensor_per_tensor_affine_stub); REGISTER_PRIVATEUSE1_DISPATCH( _fused_sdp_choice_stub, &wrapper__fused_sdp_choice); // LITERALINCLUDE END: STUB DEFAULT

其中wrapper_abs_stub只是把TensorIteratorBase转发给abs_kernel。从源码结构看,这个注册宏(定义于 DispatchStub.h)通过静态对象在程序启动时把 kernel 函数指针登记进abs_stubPrivateUse1分发槽位。

另外值得注意的是 OpenRegExtra.cpp 中一条易踩的注释:abs_stub只有在abs.out也注册到PrivateUse1之后才能真正工作,因为abs.default被设计为直接重定向到abs.out,后者再调用abs_stub。这提示我们在用 STUB 路径时,要留意 default/out 等变体之间的重定向链。

Advanced 三:Custom Operators(自定义算子)

除内置算子外,在特定场景为加速器编写自定义算子以优化性能也非常常见,通常分三类:

  1. Forward-only(仅前向)
  2. Forward and backward,分开注册
  3. Forward and backward,使用torch.autograd.Function实现

下面以最简单的"仅前向"路径为例展开。

4.1 第一步:定义 Schema

使用TORCH_LIBRARY定义算子签名(参见 OpenRegExtra.cpp):

// LITERALINCLUDE START: CUSTOM OPERATOR SCHEMA TORCH_LIBRARY(openreg, m) { m.def("custom_abs(Tensor input)-> Tensor"); } // LITERALINCLUDE END: CUSTOM OPERATOR SCHEMA

该 schema 的含义为:

  • 命名空间:openreg
  • 函数名:custom_abs
  • 输入参数:Tensor,名称为input
  • 返回类型:Tensor

4.2 第二步:注册算子

TORCH_LIBRARY_IMPLwrapper_custom_abs注册到custom_absPrivateUse1(参见 OpenRegExtra.cpp):

// LITERALINCLUDE START: CUSTOM OPERATOR DEFAULT TORCH_LIBRARY_IMPL(openreg, PrivateUse1, m) { m.impl("custom_abs", &wrapper_custom_abs); } // LITERALINCLUDE END: CUSTOM OPERATOR DEFAULT

这里有一个 Autograd 层面的重要细节:由于 PyTorch 中 Autograd 始终处于激活状态,即使只需要前向计算,Dispatcher 也会默认去寻找并执行对应的 backward 实现(即发生 fallthrough)。幸运的是,PyTorch 已为PrivateUse1实现了通用的Autograd Fallback:如果只涉及前向计算,它等效于 fallthrough 操作,选择下一个 DispatchKey 继续计算;如果涉及反向计算,则抛出错误。因此"仅前向"自定义算子无需额外处理反向注册。

4.3 第三步(可选但图模式必需):注册 Meta

自定义算子还需要注册Meta实现,torch.compiletorch.export等图模式依赖它来推断形状与元数据。PyTorch 支持在 C++ 和 Python 中注册 Meta,由于 Python 写法更简洁,参考实现选择 Python(参见 meta.py):

# LITERALINCLUDE START: CUSTOM OPERATOR META lib = torch.library.Library("openreg", "IMPL", "Meta") # noqa: SCOPED_LIBRARY @torch.library.impl(lib, "custom_abs") def custom_abs(self): return torch.empty_like(self) # LITERALINCLUDE END: CUSTOM OPERATOR META

可以看到,Python 端通过torch.library.Library(namespace, "IMPL", "Meta")构造库句柄,再用@torch.library.impl装饰器完成注册——与 C++ 端的TORCH_LIBRARY_IMPL语义一一对应,但开发体验更友好。

Tools:Dispatch 调试命令与环境变量

PyTorch 算子注册方式多样、场景众多,社区因此提供了一批工具来帮助理解底层原理与定位问题。

5.1 torch._C.dispatch* 命令

围绕 Dispatch 功能,PyTorch 提供了一组torch._C._dispatch_前缀的接口,可用如下命令查询全部相关接口:

python -c 'import torch; print("\n".join([x for x in dir(torch._C) if x.startswith("_dispatch_")]))'

典型输出包括:

... _dispatch_dump _dispatch_dump_table _dispatch_has_kernel _dispatch_has_kernel_for_any_dispatch_key _dispatch_has_kernel_for_dispatch_key _dispatch_isTensorSubclassLike _dispatch_is_alias_key _dispatch_is_included_in_alias _dispatch_is_main_interpreter _dispatch_kernel_for_dispatch_key_is_fallthrough _dispatch_key_for_device _dispatch_key_name _dispatch_key_parse _dispatch_key_set ...

几个最常用的命令:

  • torch._C._dispatch_key_set:显示当前 Tensor 的 DispatchKey 集合,优先级自左向右递增。

    >>> import torch >>> a = torch.randn(3,3,device="cuda") >>> torch._C._dispatch_key_set(a) 'DispatchKeySet(CUDA, ADInplaceOrView, AutogradCUDA, AutocastCUDA)'
  • torch._C._dispatch_dump_table:查询给定算子在各 Dispatch Key 上的支持情况,便于定位对应实现代码:

    >>> import torch >>> print(torch._C._dispatch_dump_table("aten::add.Tensor")) >>> ... CPU: registered at ./build/aten/src/ATen/RegisterCPU_0.cpp:1309 [kernel] CUDA: registered at ./build/aten/src/ATen/RegisterCUDA_0.cpp:2420 [kernel] HIP: registered at ./build/aten/src/ATen/RegisterCompositeExplicitAutogradNonFunctional_0.cpp:1373 [default backend kernel] MPS: registered at ./build/aten/src/ATen/RegisterCompositeExplicitAutogradNonFunctional_0.cpp:1373 [default backend kernel] ... PrivateUse1: registered at ./build/aten/src/ATen/RegisterCompositeExplicitAutogradNonFunctional_0.cpp:1373 [default backend kernel] ...

    通过它就能查询aten::add.Tensor在其他平台上的实现位置,从而在源码层面完整追踪算子的调用过程。对新加速器开发者来说,这是检查"我的PrivateUse1实现是否真的挂上了、是否被正确解析"的第一工具。

5.2 环境变量:TORCH_SHOW_DISPATCH_TRACE

PyTorch 还提供了若干 dispatcher 相关的环境变量,其中TORCH_SHOW_DISPATCH_TRACE可显示 PyTorch 执行期间详细的内部 Dispatch Key 调度轨迹:

export TORCH_SHOW_DISPATCH_TRACE=1
>>> import torch >>> a = torch.randn(3,3) [call] op=[aten::randn], key=[BackendSelect] [redispatch] op=[aten::randn], key=[CPU] [call] op=[aten::empty.memory_format], key=[BackendSelect] [redispatch] op=[aten::empty.memory_format], key=[CPU] [call] op=[aten::normal_], key=[CPU]

从输出中可以清楚看到 Python 级算子在 PyTorch 底层实际调用了哪些算子,包括算子名、调用层级以及对应的Dispatch Key。例如上例中一次randn就展开了empty.memory_format(内存分配)与normal_(填充随机值)两次内部调用——这正是我们 Step 1 中优先注册empty.memory_format的原因:几乎所有张量创建路径最终都会走到它。

小结:新加速器的算子适配路线图

把本文内容串联起来,一个新加速器的算子适配可按如下路线推进:

  1. 先注册最小算子集empty.memory_formatempty_stridedas_stridedviewresize__copy_from等(依据 Minimal.cpp),保证张量创建、视图与跨设备拷贝全部可用;
  2. 开启全局 Fallback:通过m.fallback+at::native::cpu_fallback让未适配算子自动回退 CPU,功能先行;
  3. 逐步收紧为选择性 Fallback:随算子适配推进,改用m.impl单算子回退或全局回退 + 黑名单的组合;
  4. 利用 STUB 降低高频算子开发成本:对DECLARE_DISPATCH声明过的算子(可用文档给出的find/grep命令检索),基于TensorIteratorBase实现 kernel,再经REGISTER_PRIVATEUSE1_DISPATCH注册;
  5. 按需定义自定义算子TORCH_LIBRARY定 schema →TORCH_LIBRARY_IMPL注册实现 →torch.library.impl注册 Meta,三步走通前向链路;
  6. 全程用工具验证torch._C._dispatch_key_set/_dispatch_dump_table核对注册状态,TORCH_SHOW_DISPATCH_TRACE=1观察真实调度路径。

以上全部代码示例均来自仓库中torch_openreg参考扩展(csrc/aten 目录),配套测试位于 tests 目录(如test_ops.pytest_autograd.py等),可作为新加速器算子注册的完整可运行范本。

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询