跳到主要内容

数据截至 (上游 commit c187ef3271d5)

02 · 算子分发:一次 torch.add 的旅程

这一章讲什么: a + b 从 Python 一路落到 CUDA kernel 的完整路径。读完你会理解 PyTorch 最核心的架构决策——dispatcher 不是「if CUDA then …」的巨型开关,而是一张按位掩码优先级的多层路由表,autograd、设备、稀疏布局、Python 子类拦截都插在同一条链上。


1. 它要解决的小问题

一个 add 要面对的维度爆炸:

  • 设备:CPU / CUDA / XPU / MPS / Meta(虚拟)……
  • 布局:dense / sparse / nested / quantized……
  • 功能层:autograd(要不要挂反向图)/ 编译器(要不要抓图)/ Python 子类(__torch_function__ 拦截)……

如果每个 kernel 自己判断「先挂图、再按设备分、再看是不是 sparse」,代码会糊成一团。PyTorch 的答案是:把每个维度编码成一个 dispatch key,按固定优先级一层层剥洋葱。


2. DispatchKey:位掩码上的优先级

DispatchKey 是一个枚举,每个值对应一个「层级」:有功能键(AutogradFunctionality)、按后端实例化的键(CPUCUDAAutogradCUDA)、别名键(CompositeImplicitAutograd)等。定义在 c10/core/DispatchKey.h:136,分类说明见同文件的 Note [DispatchKey Classification](c10/core/DispatchKey.h:114)。

优先级的实现极其朴素:每个 key 占 64 位掩码里的一位,位越高越先被处理。源码注释(c10/core/DispatchKey.h:104-112):

Higher bit indexes get handled by dispatching first (because we "count leading zeros" when we extract the highest priority dispatch key.)

每个 Tensor 在 TensorImpl::key_set_c10/core/TensorImpl.h:3045)里带着自己的 key 集合。一次调用的 dispatch key 就是把所有 Tensor 参数的 key_set 并起来再取最高位(computeDispatchKeySetaten/src/ATen/core/dispatch/DispatchKeyExtractor.h:24)。

一个典型 CUDA 张量在训练模式下,从高位到低位大致会经过:

AutogradCUDA → 先挂反向图,再往下
CUDA → 真正的 CUDA kernel
(若无注册则落 fallback / CompositeImplicitAutograd 复合实现)

3. Dispatcher::call:查表 + 调 kernel

调用的主入口是 Dispatcher::callaten/src/ATen/core/dispatch/Dispatcher.h:784)。它的快路径只有三步:

// aten/src/ATen/core/dispatch/Dispatcher.h:786-798(节选)
auto dispatchKeySet =
op.operatorDef_->op.dispatchKeyExtractor()
.getDispatchKeySetUnboxed<Args...>(args...); // ① 从参数算 key 集合
...
const KernelFunction& kernel =
op.operatorDef_->op.lookup(dispatchKeySet); // ② 查 per-op kernel 表

第 ② 步的 lookup 在每个算子的 OperatorEntry 里——它持有一张 dispatchTable_ 数组(aten/src/ATen/core/dispatch/OperatorEntry.h:237),lookup(DispatchKeySet)aten/src/ATen/core/dispatch/OperatorEntry.h:182)用 key 集合算出数组下标,直接取到 kernel 函数指针。

所以 dispatch 的本质是:算掩码 → 数前导零 → 数组下标取函数指针。 没有字符串比较、没有虚函数链,这就是它够快的原因。


4. 一张 yaml 表驱动一切

算子的「身份证」统一写在 aten/src/ATen/native/native_functions.yaml。以 add.Tensor 为例(aten/src/ATen/native/native_functions.yaml:542-551):

- func: add.Tensor(Tensor self, Tensor other, *, Scalar alpha=1) -> Tensor
device_check: NoCheck # TensorIterator
structured_delegate: add.out
variants: function, method
dispatch:
SparseCPU, SparseCUDA, ...: add_sparse
MkldnnCPU: mkldnn_add
NestedTensorCPU, ...: NestedTensor_add_Tensor
tags: [core, pointwise]

构建期由 torchgen/tools/autograd 的代码生成器读这张表,批量生成:

  • C++ 前端 API(at::add 等);
  • Python 绑定(torch.add、方法形式 a.add);
  • dispatcher 注册代码(普通条目走默认的 structured/TensorIterator 路径,dispatch: 下显式列出的 key 各注册各的 kernel);
  • autograd 的挂图层(VariableType,见第 3 章)。

一个 schema 改一处,全栈同步更新——这是 PyTorch 维持约 2500 个算子还不散架的关键。


5. 原理演示:dispatcher 在做什么

用示意代码复刻「剥洋葱」(示意,非源码):

# 示意,非源码
def dispatch(op, *args):
keyset = union_of(t.key_set for t in tensor_args(args))
key = highest_priority(keyset) # 数前导零取最高位
kernel = op.table[key] # per-op kernel 表
return kernel(op, current_keyset=keyset, *args)
# 层内干完活后,可以 redispatch:把当前 key 剥掉,再走一遍上式

真实世界里「层内干活 + 剥掉自己再分发」叫 redispatch:比如 autograd 层挂完 grad_fn 后,把 AutogradCUDA 从 key 集合里去掉,重新进 dispatcher,这次最高位就是 CUDA,于是落到真 kernel。这个模式在 dispatcher 接口上就是 redispatchaten/src/ATen/core/dispatch/Dispatcher.h:196-201)。


6. 真 kernel 在哪:stub 的两级跳

以乘法的反向视角看 sub_out 的实现(aten/src/ATen/native/BinaryOps.cpp:434-438):

TORCH_IMPL_FUNC(sub_out) (...)
{
add_stub(device_type(), *this, -alpha); // 按设备再跳一次
...
}

这里出现了第二级跳转:add_stub 是一张按设备类型索引的函数指针表(DECLARE/DEFINE_DISPATCH 机制)。各后端在编译单元里把自己的 kernel 填进去:

  • CPU:REGISTER_DISPATCH(mul_stub, &mul_kernel)aten/src/ATen/native/cpu/BinaryOpsKernel.cpp:1448
  • CUDA:REGISTER_DISPATCH(mul_stub, &mul_kernel_cuda)aten/src/ATen/native/cuda/BinaryMulKernel.cu:46

为什么要这一层?因为 dispatcher 的表是运行期全局的,而这些逐元素 kernel 想按「当前对象在哪个设备」静态、零开销地选实现(inferred:架构意图;代码上可见的是 stub 表按 device_type() 索引)。


7. Python 侧:OpOverload 与自定义注册

Python 里 torch.ops.aten.add.Tensor 拿到的是 OpOverloadtorch/_ops.py:837)。它的 __call__torch/_ops.py:910)在没有 Python 层拦截时直接进 C++ 分发;OpOverloadPackettorch/_ops.py:1237)则按签名挑具体 overload。

想注册自己的算子或给现有算子挂自定义实现,入口是 torch.library.Librarytorch/library.py:212):

  • define 写 schema(torch/library.py:272);
  • impl 把某个 Python/CUDA 实现绑到某个 dispatch key(torch/library.py:441)。

这条 Python API 最终落到和 ATen 内部完全相同的 dispatcher 注册表——自定义算子和内建算子是平权的,这就是 torch.library 生态(以及各种编译后端)能插进来的原因。


8. 坑与边界

  • key 的优先级顺序是全局约定,自定义 key 插错位置会出现「autograd 没挂上」或「子类拦截失效」这类玄学问题;官方为此在 c10/core/DispatchKey.h 写了大段分类 note,改 key 前必读。
  • fallback 与复合实现的迷雾:一个 key 没注册 kernel 时会落 backend fallback 或 CompositeImplicitAutograd,调试 dispatch 问题时可开 TORCH_SHOW_DISPATCH_TRACEDispatcher.h:790-796 有对应 trace 代码路径)。
  • TensorIterator 语义别绕开:逐元素算子的广播、类型提升在 meta 函数里完成,手写 out-of-tree kernel 若不经 TensorIterator,很容易漏掉 broadcast 边界。
  • 看不出来:各 dispatch key 在当前 build 里确切的位序——它由 torchgen/model.pyDispatchKey.h 共同决定且随版本演进,以本 commit 的 DispatchKey.h 注释为准。

9. 本章代码地图

主题文件路径符号名
dispatch key 枚举c10/core/DispatchKey.hDispatchKeyBackendSelect
key 集合与运算c10/core/DispatchKeySet.haten/src/ATen/core/dispatch/DispatchKeyExtractor.hDispatchKeySetcomputeDispatchKeySet
分发器aten/src/ATen/core/dispatch/Dispatcher.hDispatcher::callredispatch
per-op kernel 表aten/src/ATen/core/dispatch/OperatorEntry.hOperatorEntry::lookupdispatchTable_
算子 schemaaten/src/ATen/native/native_functions.yamladd.Tensor 条目(:542)
structured kernel 示例aten/src/ATen/native/BinaryOps.cppTORCH_IMPL_FUNC(sub_out)
设备级 stub 注册aten/src/ATen/native/cpu/BinaryOpsKernel.cppaten/src/ATen/native/cuda/BinaryMulKernel.cuREGISTER_DISPATCH(mul_stub, …)
Python 算子对象torch/_ops.pyOpOverload.__call__OpOverloadPacket
自定义算子 APItorch/library.pyLibrary.defineLibrary.impl