首页
/ PyTorch AT_DISPATCH_V2 宏迁移实战:从 AT_DISPATCH_* 旧宏到新 Dispatch v2 API 的完整转换指南

PyTorch AT_DISPATCH_V2 宏迁移实战:从 AT_DISPATCH_* 旧宏到新 Dispatch v2 API 的完整转换指南

2026-09-06 11:19:32作者:秋泉律Samson

PyTorch 的 ATen 层正在用一套新的类型分发宏 AT_DISPATCH_V2 逐步替代历史上以 AT_DISPATCH_ALL_TYPES_AND3AT_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 就要换一个宏名(AND2AND3AND4……),且类型组合是隐式写在宏名里的。在 aten/src/ATen/Dispatch.h 中可以确认旧宏家族的确以这种 arity 编号方式大量存在,例如 AT_DISPATCH_FLOATING_TYPES_AND2/3/4/5AT_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 一致):

  1. 参数重排scalar_typename 提到最前,随后是 lambda,最后才是类型列表;
  2. lambda 必须用 AT_WRAP 包装:防止 lambda 内部的逗号被宏解析器误认为参数分隔符;
  3. 类型组用 AT_EXPAND 展开:如 AT_EXPAND(AT_ALL_TYPES),替代旧宏的隐式展开;
  4. 逐个类型追加在类型组之后kHalfkBFloat16 等原样列出,不要再加 AT_EXPAND
  5. 补 include:在文件头部其他 Dispatch 头文件旁加上 #include <ATen/Dispatch_v2.h>

关于 AT_WRAPtorch/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_e5m2Float8_e5m2fnuzFloat8_e4m3fnFloat8_e4m3fnuzFloat8_e8m0fnu),而 AT_INTEGRAL_TYPESByte, 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_V2AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES),即整型加上 UInt16/UInt32/UInt64
  • AT_ALL_TYPES_AND_COMPLEXAT_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_SWITCHAT_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,逐个类型 kComplexHalfkHalf 直接追加在末尾——这正是 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);

无捕获 lambdaAT_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_TMPLtorch/headeronly/core/Dispatch_v2.h)。后者的机制是经典的"计数参数"宏技巧:

  1. AT_NUM_ARGS(...) 通过一个 60 项的递减数字列表统计用户传入了多少个 dtype;
  2. AT_CONCAT(THO_AP, AT_NUM_ARGS(...)) 拼接出 THO_AP1THO_AP60 中对应 arity 的手写宏,把每个类型逐个展开为 DISPATCH_CASE(type, BODY)
  3. 若类型数量超出已生成的 60 个上限,拼接会失败并产生晦涩报错。aten/src/ATen/Dispatch_v2.hstatic_assert(static_cast<int>(c10::ScalarType::NumOptions) < 60) 在编译期兜底这条约束。

文件里还保留了再生成这些宏的 Python 片段(aten/src/ATen/Dispatch_v2.h#if 0 块中的循环脚本)——若要提升 arity 上限,按注释说明需要重新生成 AT_AP1AT_AP60 系列宏。AT_EXPAND(X) Xtorch/headeronly/core/Dispatch_v2.h)则是控制宏展开时机的辅助宏,保证类型组在正确阶段被展开成完整的枚举参数列表。

仓库中的真实使用示例

v2 API 已经落地到不少 ATen 内核文件中,可直接作为转换后的参考样板:

迁移工作流与注意事项

按技能文档建议的完整工作流:

  1. 通读目标文件,找出所有 AT_DISPATCH_* 旧宏使用点;
  2. 若缺少 #include <ATen/Dispatch_v2.h> 则添加;
  3. 对每个宏依次执行:识别模式 → 提取 dtype 表达式、调试名字符串、lambda 与额外类型 → 映射基础类型组 → 构造 AT_DISPATCH_V2 调用;
  4. 逐项对照 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_AP1AT_AP60 宏限制,超限会报编译错误。

掌握以上规则后,你可以把任意旧式 AT_DISPATCH_*_AND{N} 调用安全地改写为 AT_DISPATCH_V2:参数重排、AT_WRAP 包裹 lambda、AT_EXPAND 展开类型组、额外类型裸列在末尾——四个动作覆盖所有场景,且转换结果可直接对照仓库中 aten/src/ATen/native/ 下已迁移的文件进行验证。

登录后查看全文
热门项目推荐
相关项目推荐