news 2026/9/7 18:54:09

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch AT_DISPATCH_V2 宏迁移实战:从 AT_DISPATCH_* 旧宏到新 Dispatch v2 API 的完整转换指南

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

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

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_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_TYPESAT_EXPAND(AT_ALL_TYPES)
FLOATING_TYPESAT_EXPAND(AT_FLOATING_TYPES)
INTEGRAL_TYPESAT_EXPAND(AT_INTEGRAL_TYPES)
COMPLEX_TYPESAT_EXPAND(AT_COMPLEX_TYPES)
ALL_TYPES_AND_COMPLEXAT_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),源码中标注 "notactuallyall 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.hinclude,因为同一文件里可能还有其他代码依赖它。从 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_TMPL(torch/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.h 用static_assert(static_cast<int>(c10::ScalarType::NumOptions) < 60)在编译期兜底这条约束。

文件里还保留了再生成这些宏的 Python 片段(aten/src/ATen/Dispatch_v2.h,#if 0块中的循环脚本)——若要提升 arity 上限,按注释说明需要重新生成AT_AP1AT_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(找到更多实例。

迁移工作流与注意事项

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

  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/下已迁移的文件进行验证。

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/7 18:50:48

MCP+Chrome+VSCode:让AI Agent自动操作网页的实战指南

1. MCP到底是什么&#xff0c;为什么AI操作页面非它不可 1.1 从“AI只会聊天”到“AI能动手干活”的关键一步 很多玩过ChatGPT、Claude或者其他大模型的人&#xff0c;慢慢会发现一个尴尬的事实&#xff1a;模型再聪明&#xff0c;你问它“帮我打开知乎&#xff0c;搜一下‘AI…

作者头像 李华
网站建设 2026/9/7 18:43:40

# Linux基础Day05:命令补充,用户管理及组账号管理,密码管理

声明&#xff1a;本文仅作学习交流使用&#xff0c;引用需标明出处。 如有谬误&#xff0c;敬请指正一、文件查看与处理核心命令 日常运维中需频繁查看、筛选、输出文件内容&#xff0c;这部分命令覆盖按行查看、分页查看、内容输出、结果写入、命令联动、内容过滤六大核心能力…

作者头像 李华
网站建设 2026/9/7 18:41:52

退租时押金扣多少总扯不清,如何导出微信聊天记录来核对

摘要 租房最和平分手的时候少&#xff0c;退租时闹得不愉快的多&#xff1a;墙面有钉眼算不算正常损耗、家电坏了是谁的责任、水电燃气物业费结到哪天、押金到底该扣多少&#xff0c;往往各执一词。而这些约定&#xff0c;当初全在微信里说过——入住时的房屋状况发过照片&…

作者头像 李华