SKILL.md
只读
名称
at-dispatch-v2
描述
在 ATen C++ 代码中将 PyTorch 的旧版 AT_DISPATCH 宏转换为全新的 AT_DISPATCH_V2 格式。适用于将 AT_DISPATCH_ALL_TYPES_AND*、AT_DISPATCH_FLOATING_TYPES* 或其他分发宏迁移至新版 v2 API 的场景。主要用于 ATen 算子内核文件、CUDA kernel 以及原生算子(native operator)实现。
AT_DISPATCH 至 AT_DISPATCH_V2 转换器
本 Skill 用于帮助将 PyTorch 的旧版 AT_DISPATCH 宏转换为新版的 AT_DISPATCH_V2 格式(定义于 aten/src/ATen/Dispatch_v2.h)。
适用场景
在以下情况使用本 Skill:
- 将
AT_DISPATCH_*宏转换为AT_DISPATCH_V2 - 迁移 ATen 内核代码以使用全新的分发 API
- 处理
aten/src/ATen/native/目录下使用了分发宏的文件 - 用户提到“AT_DISPATCH”、“dispatch v2”、“Dispatch_v2.h”或宏转换等关键词
快速对照
旧格式:
AT_DISPATCH_ALL_TYPES_AND3(kBFloat16, kHalf, kBool, dtype, "kernel_name", [&]() {
// lambda 函数体
});
新格式:
AT_DISPATCH_V2(dtype, "kernel_name", AT_WRAP([&]() {
// lambda 函数体
}), AT_EXPAND(AT_ALL_TYPES), kBFloat16, kHalf, kBool);
核心转换步骤
- 调整参数顺序:先传
scalar_type和name,接着是 lambda 函数,最后才是数据类型 - 包裹 Lambda 函数:使用
AT_WRAP(lambda)处理内部可能包含的逗号 - 展开类型组:使用
AT_EXPAND(AT_ALL_TYPES)代替原有的隐式展开 - 列出单独类型:在展开的类型组之后追加单独指定的类型(如 kHalf、kBFloat16 等)
- 添加头文件引用:在其他 Dispatch 头文件附近添加
#include <ATen/Dispatch_v2.h>
操作指南
第 1 步:添加 Dispatch_v2.h 引用
在已有的 #include <ATen/Dispatch.h> 附近添加 v2 版头文件:
#include <ATen/Dispatch.h>
#include <ATen/Dispatch_v2.h>
暂且保留原有的 Dispatch.h 引用(其他代码可能仍需使用)。
第 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)
第 3 步:映射旧宏到类型组
确定基础类型对应的类型组宏:
| 旧宏基础名 | 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_EXPAND() 条目:
// 旧格式:AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2(...)
// 新格式:AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_COMPLEX_TYPES), type1, type2
第 4 步:提取单独类型
从 AT_DISPATCH_*_AND2(type1, type2, ...) 或 AT_DISPATCH_*_AND3(type1, type2, type3, ...) 中提取出单独指定的数据类型(type1, type2 等)。
这些类型将作为尾部参数跟在类型组之后:
AT_DISPATCH_V2(..., AT_EXPAND(AT_ALL_TYPES), kBFloat16, kHalf, kBool)
^^^^^^^^^^^^^^^^^^^^^^^^
从 AND3 提取出的单独类型
第 5 步:转换为 AT_DISPATCH_V2
按照以下结构进行转换:
结构模式:
AT_DISPATCH_V2(
scalar_type, // 第 1 个:dtype 表达式
"name", // 第 2 个:调试字符串/内核名称
AT_WRAP(lambda), // 第 3 个:用 AT_WRAP 包裹的 lambda 函数
type_groups, // 第 4+ 个:带有 AT_EXPAND() 的类型组
individual_types // 最后:单独类型
)
转换示例:
// 转换前
AT_DISPATCH_ALL_TYPES_AND3(
kBFloat16, kHalf, kBool,
iter.dtype(),
"min_values_cuda",
[&]() {
min_values_kernel_cuda_impl<scalar_t>(iter);
}
);
// 转换后
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
);
第 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) // 内部包含逗号!
);
}),
AT_EXPAND(AT_ALL_TYPES)
);
第 7 步:验证转换结果
检查以下项目:
- [ ]
AT_WRAP()包裹了完整的 lambda 函数 - [ ] 类型组使用了
AT_EXPAND() - [ ] 单独类型未添加
AT_EXPAND()(直接写kBFloat16,而非AT_EXPAND(kBFloat16)) - [ ] 参数顺序准确:scalar_type、name、lambda、types
- [ ] 已添加头文件包含:
#include <ATen/Dispatch_v2.h>
类型组参考手册
可用的类型组宏(搭配 AT_EXPAND() 使用):
AT_INTEGRAL_TYPES // kByte, kChar, kInt, kLong, kShort
AT_FLOATING_TYPES // kDouble, kFloat
AT_COMPLEX_TYPES // kComplexDouble, kComplexFloat
AT_QINT_TYPES // kQInt8, kQUInt8, kQInt32
AT_ALL_TYPES // INTEGRAL_TYPES + FLOATING_TYPES
AT_ALL_TYPES_AND_COMPLEX // ALL_TYPES + COMPLEX_TYPES
AT_INTEGRAL_TYPES_V2 // INTEGRAL_TYPES + 无符号类型
AT_BAREBONES_UNSIGNED_TYPES // kUInt16, kUInt32, kUInt64
AT_FLOAT8_TYPES // Float8 变体类型
常见转换模式
模式 1:AT_DISPATCH_ALL_TYPES_AND2
// 转换前
AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBFloat16, dtype, "op", [&]() {
kernel<scalar_t>(data);
});
// 转换后
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>(data);
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);
模式 2:AT_DISPATCH_FLOATING_TYPES_AND3
// 转换前
AT_DISPATCH_FLOATING_TYPES_AND3(kHalf, kBFloat16, kFloat8_e4m3fn,
tensor.scalar_type(), "float_op", [&] {
process<scalar_t>(tensor);
});
// 转换后
AT_DISPATCH_V2(tensor.scalar_type(), "float_op", AT_WRAP([&] {
process<scalar_t>(tensor);
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf, kBFloat16, kFloat8_e4m3fn);
模式 3:AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2
// 转换前
AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2(
kComplexHalf, kHalf,
self.scalar_type(),
"complex_op",
[&] {
result = compute<scalar_t>(self);
}
);
// 转换后
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
);
边缘情况处理
场景 1:无额外追加类型(较少出现)
// 转换前
AT_DISPATCH_ALL_TYPES(dtype, "op", [&]() { kernel<scalar_t>(); });
// 转换后
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES));
场景 2:包含多个单独类型(AND4, AND5 等)
// 转换前
AT_DISPATCH_FLOATING_TYPES_AND4(kHalf, kBFloat16, kFloat8_e4m3fn, kFloat8_e5m2,
dtype, "float8_op", [&]() { kernel<scalar_t>(); });
// 转换后
AT_DISPATCH_V2(dtype, "float8_op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf, kBFloat16, kFloat8_e4m3fn, kFloat8_e5m2);
场景 3:无捕获列表的 Lambda 函数
// 转换前
AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBool, dtype, "op", []() {
static_kernel<scalar_t>();
});
// 转换后
AT_DISPATCH_V2(dtype, "op", AT_WRAP([]() {
static_kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBool);
AT_DISPATCH_V2 的优势
- 宏名称与参数数量解耦:无需再针对 AND2、AND3、AND4 使用不同的宏名称
- 类型组合更灵活:可通过
AT_EXPAND()自由搭积木式组合不同类型组 - 扩展性更佳:轻松添加更多类型,不受宏参数数量上限限制
- 语义更清晰:类型组显式表达,不再隐含于宏名称中
重要注意事项
- 务必保留
#include <ATen/Dispatch.h>—— 其他未改造代码仍需要它 AT_WRAP()是强制要求的 —— 用于防止 lambda 函数内的逗号被宏解析错乱- 类型组必须使用
AT_EXPAND(),而单独数据类型不需要 - v2 版 API 位于
aten/src/ATen/Dispatch_v2.h—— 可查阅该文件获取完整文档 - 头文件中还包含了重新生成该宏实现的 Python 脚本
工作流程
当收到转换 AT_DISPATCH 宏的需求时:
- 通读目标文件,识别所有
AT_DISPATCH使用位置 - 若尚未引用
#include <ATen/Dispatch_v2.h>,则添加该头文件 - 针对每个分发宏:
- 识别模式并提取相关组件
- 映射基础类型组
- 提取单独的数据类型
- 构造对应的
AT_DISPATCH_V2调用 - 使用 Edit 工具在代码中应用修改
- 向用户展示转换完成后的完整文件
- 简要说明改动内容
请不要编译或测试代码 —— 专注于确保转换过程的准确无误。






