数据截至 (上游 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)、按后端实例化的键(CPU、CUDA、AutogradCUDA)、别名键(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 并起来再取最高位(computeDispatchKeySet,aten/src/ATen/core/dispatch/DispatchKeyExtractor.h:24)。
一个典型 CUDA 张量在训练模式下,从高位到低位大致会经过:
AutogradCUDA → 先挂反向图,再往下
CUDA → 真正的 CUDA kernel
(若无注册则落 fallback / CompositeImplicitAutograd 复合实现)
3. Dispatcher::call:查表 + 调 kernel
调用的主入 口是 Dispatcher::call(aten/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 接口上就是 redispatch(aten/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 拿到的是 OpOverload(torch/_ops.py:837)。它的 __call__(torch/_ops.py:910)在没有 Python 层拦截时直接进 C++ 分发;OpOverloadPacket(torch/_ops.py:1237)则按签名挑具体 overload。
想注册自己 的算子或给现有算子挂自定义实现,入口是 torch.library.Library(torch/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_TRACE(Dispatcher.h:790-796有对应 trace 代码路径)。 - TensorIterator 语义别绕开:逐元素算子的广播、类型提升在 meta 函数里完成,手写 out-of-tree kernel 若不经 TensorIterator,很容易漏掉 broadcast 边界。
- 看不出来:各 dispatch key 在当前 build 里确切的位序——它由
torchgen/model.py与DispatchKey.h共同决定且随版本演进,以本 commit 的DispatchKey.h注释为准。
9. 本章代 码地图
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| dispatch key 枚举 | c10/core/DispatchKey.h | DispatchKey、BackendSelect |
| key 集合与运算 | c10/core/DispatchKeySet.h、aten/src/ATen/core/dispatch/DispatchKeyExtractor.h | DispatchKeySet、computeDispatchKeySet |
| 分发器 | aten/src/ATen/core/dispatch/Dispatcher.h | Dispatcher::call、redispatch |
| per-op kernel 表 | aten/src/ATen/core/dispatch/OperatorEntry.h | OperatorEntry::lookup、dispatchTable_ |
| 算子 schema | aten/src/ATen/native/native_functions.yaml | add.Tensor 条目(:542) |
| structured kernel 示例 | aten/src/ATen/native/BinaryOps.cpp | TORCH_IMPL_FUNC(sub_out) |
| 设备级 stub 注册 | aten/src/ATen/native/cpu/BinaryOpsKernel.cpp、aten/src/ATen/native/cuda/BinaryMulKernel.cu | REGISTER_DISPATCH(mul_stub, …) |
| Python 算子对象 | torch/_ops.py | OpOverload.__call__、OpOverloadPacket |
| 自定义算子 API | torch/library.py | Library.define、Library.impl |