ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

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

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/pytorchPyTorch 的 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_V2AT_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_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_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, kHalfv2 头文件中实际可用的类型组宏定义在 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)源码中标注 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/UInt64AT_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_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_implscalar_t(iter); } ); // AFTER AT_DISPATCH_V2( iter.dtype(), min_values_cuda, AT_WRAP([]() { min_values_kernel_cuda_implscalar_t(iter); }), AT_EXPAND(AT_ALL_TYPES), kBFloat16, kHalf, kBool );Step 6处理多行/含逗号的 lambdalambda 内部有逗号时AT_WRAP是必需的AT_DISPATCH_V2( dtype, complex_kernel, AT_WRAP([]() { gpu_reduce_kernelscalar_t, scalar_t( iter, MinOpsscalar_t{}, thrust::pairscalar_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, []() { kernelscalar_t(data); }); // After AT_DISPATCH_V2(dtype, op, AT_WRAP([]() { kernelscalar_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, [] { processscalar_t(tensor); }); // After AT_DISPATCH_V2(tensor.scalar_type(), float_op, AT_WRAP([] { processscalar_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 computescalar_t(self); } ); // After AT_DISPATCH_V2( self.scalar_type(), complex_op, AT_WRAP([] { result computescalar_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, []() { kernelscalar_t(); }); // After AT_DISPATCH_V2(dtype, op, AT_WRAP([]() { kernelscalar_t(); }), AT_EXPAND(AT_ALL_TYPES));大量额外类型AND4/AND5——v2 的一个优势是这种场景不再受宏名限制// Before AT_DISPATCH_FLOATING_TYPES_AND4(kHalf, kBFloat16, kFloat8_e4m3fn, kFloat8_e5m2, dtype, float8_op, []() { kernelscalar_t(); }); // After AT_DISPATCH_V2(dtype, float8_op, AT_WRAP([]() { kernelscalar_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_castc10::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。后者的机制是经典的计数参数宏技巧AT_NUM_ARGS(...)通过一个 60 项的递减数字列表统计用户传入了多少个 dtypeAT_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_castint(c10::ScalarType::NumOptions) 60)在编译期兜底这条约束。文件里还保留了再生成这些宏的 Python 片段aten/src/ATen/Dispatch_v2.h#if 0块中的循环脚本——若要提升 arity 上限按注释说明需要重新生成AT_AP1…AT_AP60系列宏。AT_EXPAND(X) Xtorch/headeronly/core/Dispatch_v2.h则是控制宏展开时机的辅助宏保证类型组在正确阶段被展开成完整的枚举参数列表。仓库中的真实使用示例v2 API 已经落地到不少 ATen 内核文件中可直接作为转换后的参考样板aten/src/ATen/native/cpu/FillKernel.cppfill_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.h60 个类型上限单次AT_DISPATCH_V2调用展开的 dtype 总数受已生成的AT_AP1–AT_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),仅供参考
返回列表