at-dispatch-v2

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)实现。

10万Star
2.9万Fork
更新于 2026/8/4
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);

核心转换步骤

  1. 调整参数顺序:先传 scalar_typename,接着是 lambda 函数,最后才是数据类型
  2. 包裹 Lambda 函数:使用 AT_WRAP(lambda) 处理内部可能包含的逗号
  3. 展开类型组:使用 AT_EXPAND(AT_ALL_TYPES) 代替原有的隐式展开
  4. 列出单独类型:在展开的类型组之后追加单独指定的类型(如 kHalf、kBFloat16 等)
  5. 添加头文件引用:在其他 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 的优势

  1. 宏名称与参数数量解耦:无需再针对 AND2、AND3、AND4 使用不同的宏名称
  2. 类型组合更灵活:可通过 AT_EXPAND() 自由搭积木式组合不同类型组
  3. 扩展性更佳:轻松添加更多类型,不受宏参数数量上限限制
  4. 语义更清晰:类型组显式表达,不再隐含于宏名称中

重要注意事项

  • 务必保留 #include <ATen/Dispatch.h> —— 其他未改造代码仍需要它
  • AT_WRAP() 是强制要求的 —— 用于防止 lambda 函数内的逗号被宏解析错乱
  • 类型组必须使用 AT_EXPAND(),而单独数据类型不需要
  • v2 版 API 位于 aten/src/ATen/Dispatch_v2.h —— 可查阅该文件获取完整文档
  • 头文件中还包含了重新生成该宏实现的 Python 脚本

工作流程

当收到转换 AT_DISPATCH 宏的需求时:

  1. 通读目标文件,识别所有 AT_DISPATCH 使用位置
  2. 若尚未引用 #include <ATen/Dispatch_v2.h>,则添加该头文件
  3. 针对每个分发宏:
    • 识别模式并提取相关组件
    • 映射基础类型组
    • 提取单独的数据类型
    • 构造对应的 AT_DISPATCH_V2 调用
    • 使用 Edit 工具在代码中应用修改
  4. 向用户展示转换完成后的完整文件
  5. 简要说明改动内容

不要编译或测试代码 —— 专注于确保转换过程的准确无误。