PyTorch AT_DISPATCH_V2 宏迁移实战:从 AT_DISPATCH_* 旧宏到新 Dispatch v2 API 的完整转换指南
PyTorch 的 ATen 层正在用一套新的类型分发宏 AT_DISPATCH_V2 逐步替代历史上以 AT_DISPATCH_ALL_TYPES_AND3、AT_DISPATCH_FLOATING_TYPES_AND2 等命名的旧宏家族。本文基于 PyTorch 仓库中的迁移技能文档 .claude/skills/at-dispatch-v2/SKILL.md,完整讲解新旧两种写法的参数差异、类型组(type group)映射关系、AT_WRAP/AT_EXPAND 等辅助宏的作用,并结合 aten/src/ATen/Dispatch_v2.h 的实现源码与实际内核代码(如 aten/src/ATen/native/cpu/FillKernel.cpp)逐条佐证,帮你在编写或移植 ATen 内核时正确使用 v2 分发 API。
为什么需要 AT_DISPATCH_V2:旧宏的痛点
ATen 内核需要根据 Tensor 的实际 dtype 实例化不同的模板特化,这个过程由 ATen/Dispatch.h 中的宏家族完成。aten/src/ATen/Dispatch.h 的注释说明了旧式用法:
AT_DISPATCH_ALL_TYPES(self.scalar_type(), "op_name", [&] {
// 'scalar_t' 在此被定义为当前 dtype
});
旧宏的核心限制在于:宏名本身编码了"基础类型组 + 额外类型个数"两个维度,因此每多一个额外 dtype 就要换一个宏名(AND2、AND3、AND4……),且类型组合是隐式写在宏名里的。在 aten/src/ATen/Dispatch.h 中可以确认旧宏家族的确以这种 arity 编号方式大量存在,例如 AT_DISPATCH_FLOATING_TYPES_AND2/3/4/5、AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2/3/4/5/6/7/8 等。
v2 API 针对这些痛点做了三项改进(见 aten/src/ATen/Dispatch_v2.h 头部注释):
- 不再需要指定 arity:无需
AND{2,3,4,...}式宏名,AT_DISPATCH_V2一个宏覆盖所有参数个数; - 类型集合可组合:相关的一组 dtype 可直接写
AT_EXPAND(AT_INTEGRAL_TYPES)这类类型组,无需逐个罗列; - 类型显式:类型组在参数列表中显式出现,而不是隐式编码在宏名中。
新旧格式对照:参数顺序与包装规则
旧格式速览
迁移技能文档给出的旧格式示例:
AT_DISPATCH_ALL_TYPES_AND3(kBFloat16, kHalf, kBool, dtype, "kernel_name", [&]() {
// lambda body
});
参数顺序是:额外类型1..n, scalar_type 表达式, 调试用名字符串, lambda。
新格式(AT_DISPATCH_V2)
AT_DISPATCH_V2(dtype, "kernel_name", AT_WRAP([&]() {
// lambda body
}), AT_EXPAND(AT_ALL_TYPES), kBFloat16, kHalf, kBool);
参数顺序发生了重排,这是转换时最容易出错的地方。AT_DISPATCH_V2 的完整签名(见 aten/src/ATen/Dispatch_v2.h)为:
AT_DISPATCH_V2(
scalar_type, // 第 1 参:dtype 表达式(如 iter.dtype())
"name", // 第 2 参:调试字符串(算子名)
AT_WRAP(lambda), // 第 3 参:用 AT_WRAP 包装的 lambda
type_groups, // 第 4 参起:类型组,需 AT_EXPAND()
individual_types // 末尾:逐个列出的额外类型
)
五个关键转换动作(与技能文档 Key transformations 一致):
- 参数重排:
scalar_type与name提到最前,随后是 lambda,最后才是类型列表; - lambda 必须用
AT_WRAP包装:防止 lambda 内部的逗号被宏解析器误认为参数分隔符; - 类型组用
AT_EXPAND展开:如AT_EXPAND(AT_ALL_TYPES),替代旧宏的隐式展开; - 逐个类型追加在类型组之后:
kHalf、kBFloat16等原样列出,不要再加AT_EXPAND; - 补 include:在文件头部其他 Dispatch 头文件旁加上
#include <ATen/Dispatch_v2.h>。
关于 AT_WRAP,torch/headeronly/core/Dispatch_v2.h 给出了定义和注释:它是一个"把可能包含内部逗号的任意表达式传递给另一个宏而不被拆散"的工具,定义即 #define AT_WRAP(...) __VA_ARGS__。而 aten/src/ATen/Dispatch_v2.h 明确提醒:"必须记住用 AT_WRAP 包装 payload body,否则 lambda 里的逗号会被错误处理"。
旧宏到 v2 类型组的映射表
转换的核心是把旧宏前缀映射为 v2 的类型组宏。映射关系如下:
| 旧宏前缀 | AT_DISPATCH_V2 类型组 |
|---|---|
ALL_TYPES |
AT_EXPAND(AT_ALL_TYPES) |
FLOATING_TYPES |
AT_EXPAND(AT_FLOATING_TYPES) |
INTEGRAL_TYPES |
AT_EXPAND(AT_INTEGRAL_TYPES) |
COMPLEX_TYPES |
AT_EXPAND(AT_COMPLEX_TYPES) |
ALL_TYPES_AND_COMPLEX |
AT_EXPAND(AT_ALL_TYPES_AND_COMPLEX) |
对"复合"旧宏(如 AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2),拆成多个 AT_EXPAND() 条目再加逐个类型:
// 旧: AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2(kComplexHalf, kHalf, ...)
// 新: AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_COMPLEX_TYPES), kComplexHalf, kHalf
v2 头文件中实际可用的类型组宏定义在 torch/headeronly/core/Dispatch_v2.h,其内容比技能文档的速查表更完整,例如 AT_FLOAT8_TYPES 实际包含 5 个 Float8 变体(Float8_e5m2、Float8_e5m2fnuz、Float8_e4m3fn、Float8_e4m3fnuz、Float8_e8m0fnu),而 AT_INTEGRAL_TYPES 是 Byte, Char, Int, Long, Short 五个无符号/有符号整型,AT_FLOATING_TYPES 仅为 Double, Float。注意 AT_ALL_TYPES 的源码定义是 AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_FLOATING_TYPES),源码中标注 "not actually all types"——它不包含 Bool、Half、Complex,这与旧 AT_DISPATCH_ALL_TYPES 的语义一致(历史原因,见 aten/src/ATen/Dispatch.h 注释)。
另外两个值得知道的组合宏:
AT_INTEGRAL_TYPES_V2:AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES),即整型加上UInt16/UInt32/UInt64;AT_ALL_TYPES_AND_COMPLEX:AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_COMPLEX_TYPES)。
逐步转换流程与完整示例
Step 1:添加头文件
在原有 #include <ATen/Dispatch.h> 旁边加上 v2 头文件:
#include <ATen/Dispatch.h>
#include <ATen/Dispatch_v2.h>
迁移期间建议保留旧的 Dispatch.h include,因为同一文件里可能还有其他代码依赖它。从 aten/src/ATen/Dispatch_v2.h 的源码也能看到,v2 头文件本身就 include 了 Dispatch.h(为了复用 AT_DISPATCH_SWITCH 和 AT_DISPATCH_CASE),所以旧 include 并不冲突。
Step 2:识别旧模式
需要转换的常见旧模式:
AT_DISPATCH_ALL_TYPES_AND{2,3,4}(type1, type2, ..., scalar_type, name, lambda)AT_DISPATCH_FLOATING_TYPES_AND{2,3}(type1, type2, ..., scalar_type, name, lambda)AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND{2,3}(type1, ..., scalar_type, name, lambda)AT_DISPATCH_FLOATING_AND_COMPLEX_TYPES_AND{2,3}(type1, ..., scalar_type, name, lambda)
Step 3~5:映射类型组、提取额外类型、构造新调用
从 AND2/AND3 的前导参数中提取逐个类型,作为类型组之后的尾部参数。技能文档给出的标准转换示例:
// BEFORE
AT_DISPATCH_ALL_TYPES_AND3(
kBFloat16, kHalf, kBool,
iter.dtype(),
"min_values_cuda",
[&]() {
min_values_kernel_cuda_impl<scalar_t>(iter);
}
);
// AFTER
AT_DISPATCH_V2(
iter.dtype(),
"min_values_cuda",
AT_WRAP([&]() {
min_values_kernel_cuda_impl<scalar_t>(iter);
}),
AT_EXPAND(AT_ALL_TYPES),
kBFloat16, kHalf, kBool
);
Step 6:处理多行/含逗号的 lambda
lambda 内部有逗号时,AT_WRAP 是必需的:
AT_DISPATCH_V2(
dtype,
"complex_kernel",
AT_WRAP([&]() {
gpu_reduce_kernel<scalar_t, scalar_t>(
iter,
MinOps<scalar_t>{},
thrust::pair<scalar_t, int64_t>(upper_bound(), 0) // lambda 内部有逗号
);
}),
AT_EXPAND(AT_ALL_TYPES)
);
Step 7:转换后自检清单
- [ ]
AT_WRAP()完整包裹了整个 lambda; - [ ] 类型组都用了
AT_EXPAND(); - [ ] 逐个类型没有加
AT_EXPAND()(写kBFloat16而不是AT_EXPAND(kBFloat16)); - [ ] 参数顺序为
scalar_type, name, lambda, types; - [ ] 已添加
#include <ATen/Dispatch_v2.h>。
常见模式转换速查
模式一:AT_DISPATCH_ALL_TYPES_AND2
// Before
AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBFloat16, dtype, "op", [&]() {
kernel<scalar_t>(data);
});
// After
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>(data);
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);
模式二:AT_DISPATCH_FLOATING_TYPES_AND3
// Before
AT_DISPATCH_FLOATING_TYPES_AND3(kHalf, kBFloat16, kFloat8_e4m3fn,
tensor.scalar_type(), "float_op", [&] {
process<scalar_t>(tensor);
});
// After
AT_DISPATCH_V2(tensor.scalar_type(), "float_op", AT_WRAP([&] {
process<scalar_t>(tensor);
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf, kBFloat16, kFloat8_e4m3fn);
模式三:AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2(复合类型组)
// Before
AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2(
kComplexHalf, kHalf,
self.scalar_type(),
"complex_op",
[&] {
result = compute<scalar_t>(self);
}
);
// After
AT_DISPATCH_V2(
self.scalar_type(),
"complex_op",
AT_WRAP([&] {
result = compute<scalar_t>(self);
}),
AT_EXPAND(AT_ALL_TYPES),
AT_EXPAND(AT_COMPLEX_TYPES),
kComplexHalf,
kHalf
);
这里两个类型组各用一次 AT_EXPAND,逐个类型 kComplexHalf、kHalf 直接追加在末尾——这正是 aten/src/ATen/Dispatch_v2.h 头部注释中给出的官方对照示例(_local_scalar_dense_cpu 的转换)。
边缘情况
无额外类型(旧宏本身不带 AND):
// Before
AT_DISPATCH_ALL_TYPES(dtype, "op", [&]() { kernel<scalar_t>(); });
// After
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES));
大量额外类型(AND4/AND5)——v2 的一个优势是这种场景不再受宏名限制:
// Before
AT_DISPATCH_FLOATING_TYPES_AND4(kHalf, kBFloat16, kFloat8_e4m3fn, kFloat8_e5m2,
dtype, "float8_op", [&]() { kernel<scalar_t>(); });
// After
AT_DISPATCH_V2(dtype, "float8_op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf, kBFloat16, kFloat8_e4m3fn, kFloat8_e5m2);
无捕获 lambda:AT_WRAP([]() {...}) 与有捕获情况写法一致,只是括号内捕获列表为空。
源码层实现原理:v2 宏到底做了什么
从源码结构看,AT_DISPATCH_V2 本身只是一个薄薄的封装(aten/src/ATen/Dispatch_v2.h):
#define AT_DISPATCH_V2(TYPE, NAME, BODY, ...) \
THO_DISPATCH_V2_TMPL( \
AT_DISPATCH_SWITCH, \
AT_DISPATCH_CASE, \
TYPE, NAME, AT_WRAP(BODY), __VA_ARGS__)
它把 AT_DISPATCH_SWITCH(生成 switch (static_cast<c10::ScalarType>(TYPE)) 的 switch 语句)和 AT_DISPATCH_CASE(生成每个 case enum_type: { using scalar_t = ...; return BODY(); } 分支,定义于 aten/src/ATen/Dispatch.h)作为参数传给通用的 THO_DISPATCH_V2_TMPL(torch/headeronly/core/Dispatch_v2.h)。后者的机制是经典的"计数参数"宏技巧:
AT_NUM_ARGS(...)通过一个 60 项的递减数字列表统计用户传入了多少个 dtype;AT_CONCAT(THO_AP, AT_NUM_ARGS(...))拼接出THO_AP1…THO_AP60中对应 arity 的手写宏,把每个类型逐个展开为DISPATCH_CASE(type, BODY);- 若类型数量超出已生成的 60 个上限,拼接会失败并产生晦涩报错。aten/src/ATen/Dispatch_v2.h 用
static_assert(static_cast<int>(c10::ScalarType::NumOptions) < 60)在编译期兜底这条约束。
文件里还保留了再生成这些宏的 Python 片段(aten/src/ATen/Dispatch_v2.h,#if 0 块中的循环脚本)——若要提升 arity 上限,按注释说明需要重新生成 AT_AP1…AT_AP60 系列宏。AT_EXPAND(X) X(torch/headeronly/core/Dispatch_v2.h)则是控制宏展开时机的辅助宏,保证类型组在正确阶段被展开成完整的枚举参数列表。
仓库中的真实使用示例
v2 API 已经落地到不少 ATen 内核文件中,可直接作为转换后的参考样板:
- aten/src/ATen/native/cpu/FillKernel.cpp:
fill_cpu使用AT_DISPATCH_V2(iter.dtype(), "fill_cpu", AT_WRAP(...), AT_EXPAND(AT_ALL_TYPES_AND_COMPLEX), kBool, AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)),演示了"类型组 + 逐个类型混合"的写法;注意非原生类型(Half、BFloat16、各 Float8)在宏之外用if/else分支处理。 - aten/src/ATen/native/Scalar.cpp:
_local_scalar_dense_cpu使用自定义类型组AT_SD_TYPES(基类类型加上AT_EXPAND(AT_FLOAT8_TYPES)),说明 v2 API 支持先#define自己的类型组组合,再整体AT_EXPAND传入。 - aten/src/ATen/native/cuda/Copy.cu、aten/src/ATen/native/cpu/CopyKernel.cpp、aten/src/ATen/native/ReduceOps.cpp 等文件也已采用
AT_DISPATCH_V2,可以检索AT_DISPATCH_V2(找到更多实例。
迁移工作流与注意事项
按技能文档建议的完整工作流:
- 通读目标文件,找出所有
AT_DISPATCH_*旧宏使用点; - 若缺少
#include <ATen/Dispatch_v2.h>则添加; - 对每个宏依次执行:识别模式 → 提取 dtype 表达式、调试名字符串、lambda 与额外类型 → 映射基础类型组 → 构造
AT_DISPATCH_V2调用; - 逐项对照 Step 7 自检清单核对转换结果。
几点必须遵守的注意事项(来自文档 Important notes 与源码事实):
- 保留
#include <ATen/Dispatch.h>:其他代码可能仍在使用旧宏与AT_DISPATCH_SWITCH/CASE基础设施; AT_WRAP()不可省略:它是 lambda 内部逗号不被宏拆解的唯一保障;- 类型组必须
AT_EXPAND(),逐个类型不要:AT_EXPAND(kBFloat16)这种写法是错误示范; - v2 API 权威定义在 aten/src/ATen/Dispatch_v2.h,遇到本文未覆盖的用法(如自定义
THO_DISPATCH_V2_TMPL派生宏)应直接查阅该文件与 torch/headeronly/core/Dispatch_v2.h; - 60 个类型上限:单次
AT_DISPATCH_V2调用展开的 dtype 总数受已生成的AT_AP1–AT_AP60宏限制,超限会报编译错误。
掌握以上规则后,你可以把任意旧式 AT_DISPATCH_*_AND{N} 调用安全地改写为 AT_DISPATCH_V2:参数重排、AT_WRAP 包裹 lambda、AT_EXPAND 展开类型组、额外类型裸列在末尾——四个动作覆盖所有场景,且转换结果可直接对照仓库中 aten/src/ATen/native/ 下已迁移的文件进行验证。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0623
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00