首页
/ PyTorch 算子类型扩展实战:用 AT_DISPATCH_V2 为 uint16/uint32/uint64 补齐无符号整数支持

PyTorch 算子类型扩展实战:用 AT_DISPATCH_V2 为 uint16/uint32/uint64 补齐无符号整数支持

2026-09-03 16:21:17作者:咎岭娴Homer

本文为 PyTorch 开发者讲解如何为算子内核(kernel)添加无符号整数类型(kUInt16kUInt32kUInt64)的分发支持:核心手法是在算子实现的 AT_DISPATCH_V2 宏中引入 AT_BAREBONES_UNSIGNED_TYPES 类型组,或把 AT_INTEGRAL_TYPES 替换为其超集 AT_INTEGRAL_TYPES_V2。读完本文,你将掌握判断何时需要转换到 V2 宏、两种改法的适用条件、多分发点的一致修改策略,以及基于 aten/src/ATen/Dispatch_v2.htorch/headeronly/core/Dispatch_v2.h 的类型组定义来验证改动正确性的完整方法。

背景:AT_DISPATCH 宏如何决定算子支持哪些 dtype

PyTorch 的算子内核在 C++ 层通过 AT_DISPATCH 宏家族做运行时类型分发:宏会根据传入的 ScalarType(运行时 dtype)展开为一段 switch 语句,把每个受支持的 dtype 映射到对应的 C++ 标量类型 scalar_t,再调用你的内核模板函数。一个 dtype 是否被支持,完全取决于你写在分发宏参数列表里的类型——不在列表中的 dtype 会直接走到默认分支并报错

新版分发宏 AT_DISPATCH_V2 定义于 aten/src/ATen/Dispatch_v2.h,它相比旧版 AT_DISPATCH_ALL_TYPES_AND2(...) 之类的宏有两个关键改进(源码注释中明确说明):

  • 不再需要按"额外类型个数"选择带后缀的宏名(AT_DISPATCH_..._AND2/AND3/...),可变参数直接写即可;
  • 支持直接传入"类型组"宏(如 AT_EXPAND(AT_INTEGRAL_TYPES)),而不用逐个列举 dtype。

其调用形态为:

AT_DISPATCH_V2(
  scalar_type,        // 运行时 ScalarType
  "debug string",     // 出错的算子名,便于报错定位
  AT_WRAP([&] {       // 必须用 AT_WRAP 包住 lambda,否则内部逗号会被误判为宏参数
    ... code to specialize with scalar_t ...
  }),
  kHalf,
  AT_EXPAND(AT_ALL_TYPES),
  ... as many type arguments as needed ...
)

这里有一个容易踩的坑:AT_WRAP 不能省略。宏参数以逗号分隔,若 lambda 体内出现逗号而不包在 AT_WRAP 里,编译器会把它们当成额外的宏实参,产生难以理解的报错。源码注释还特别指出:如果传入的 dtype 数量超过宏支持的上限,会出现"晦涩的报错"(源于尝试把 AT_AP 与非数字拼接)——从 aten/src/ATen/Dispatch_v2.h 可以看到有一条 static_assert(static_cast<int>(c10::ScalarType::NumOptions) < 60) 约束,生成的展开宏 AT_AP1AT_AP60 最多支持 60 个类型实参。

类型组参考:谁包含了无符号类型

类型组的定义集中在 torch/headeronly/core/Dispatch_v2.h。与无符号整数支持直接相关的宏定义如下:

// 有符号整数组
#define AT_INTEGRAL_TYPES                                                      \
  torch::headeronly::ScalarType::Byte, torch::headeronly::ScalarType::Char,    \
      torch::headeronly::ScalarType::Int, torch::headeronly::ScalarType::Long, \
      torch::headeronly::ScalarType::Short

// 无符号整数组(barebones,"最小集")
#define AT_BAREBONES_UNSIGNED_TYPES          \
  torch::headeronly::ScalarType::UInt16,     \
      torch::headeronly::ScalarType::UInt32, \
      torch::headeronly::ScalarType::UInt64

// V2 整数组 = 有符号整数 + 无符号整数
#define AT_INTEGRAL_TYPES_V2 \
  AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)

// 注意:AT_ALL_TYPES 并不包含无符号类型
#define AT_ALL_TYPES AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_FLOATING_TYPES)

三者的包含关系可以概括为:

AT_INTEGRAL_TYPES            // kByte, kChar, kInt, kLong, kShort
AT_BAREBONES_UNSIGNED_TYPES  // kUInt16, kUInt32, kUInt64
AT_INTEGRAL_TYPES_V2         // INTEGRAL_TYPES + BAREBONES_UNSIGNED_TYPES

这里有两个要点需要特别注意:

  1. AT_ALL_TYPES 不含无符号类型。源码中 AT_ALL_TYPES 的展开仅为 AT_INTEGRAL_TYPES + AT_FLOATING_TYPES,因此只写了 AT_EXPAND(AT_ALL_TYPES) 的分发并不支持 uint16/32/64,必须显式追加 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES) 才生效。
  2. kByte(uint8)不属于"barebones"无符号组kByte 已经包含在 AT_INTEGRAL_TYPES 中;AT_BAREBONES_UNSIGNED_TYPES 仅指 kUInt16kUInt32kUInt64 三个类型。另外 torch/headeronly/core/Dispatch_v2.h 中还有一个 AT_OPAQUE_TYPES(Byte、UInt16、UInt32、UInt64、ComplexDouble),用于"把张量当不透明字节块处理"的场景,与本文的整数算子支持不是一回事。

AT_EXPAND(X) 本身是一个恒等展开宏(#define AT_EXPAND(X) X,见 torch/headeronly/core/Dispatch_v2.h),作用是强制先展开类型组宏、再参与外层参数计数,所以类型组必须用 AT_EXPAND() 包裹,直接裸写宏名会导致参数计数错乱。

判断起点:文件是否已在用 AT_DISPATCH_V2

改动前第一步永远是确认目标文件使用的分发宏版本:

  • 如果还在用旧版 AT_DISPATCH 系列宏(例如 AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBFloat16, dtype, "op", ...)),需要先把该分发点改写为 AT_DISPATCH_V2 形态,然后再加无符号类型。改写的对应关系是固定的:旧宏的"类型组 + 额外类型"逐一映射为 V2 的实参列表。
  • 如果已经在用 AT_DISPATCH_V2,直接进入下一步分析。

旧到新的一次完整改法示例(同时演示了加 uint 的最终形态):

// Before(旧版宏,不支持无符号类型)
AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBFloat16, dtype, "op", [&]() {
  kernel<scalar_t>();
});

// After v2 conversion(先转换到 V2)
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);

// After adding uint support(再追加无符号类型组)
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES), kHalf, kBFloat16);

转换 V2 之后,第二步是识别当前分发宏里的类型覆盖情况,常见模式有三种:

AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  // body
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);
    ^^^^^^^^^^^^^^^^^^^^^^^^^
    当前类型覆盖
  • AT_EXPAND(AT_ALL_TYPES) → 有符号整数 + 浮点,不含无符号 16/32/64;
  • AT_EXPAND(AT_INTEGRAL_TYPES) → 仅有符号整数;
  • AT_EXPAND(AT_FLOATING_TYPES) → 仅浮点类型。

两种添加无符号类型的改法

方法一:显式追加 AT_BAREBONES_UNSIGNED_TYPES

适用于任何已有整数或全类型覆盖的场景,语义最直白——"在现有类型上额外增加无符号类型":

// Before
AT_DISPATCH_V2(
    dtype,
    "min_values_cuda",
    AT_WRAP([&]() {
      kernel_impl<scalar_t>(iter);
    }),
    AT_EXPAND(AT_ALL_TYPES),
    kBFloat16, kHalf, kBool
);

// After(在类型列表中追加无符号类型组)
AT_DISPATCH_V2(
    dtype,
    "min_values_cuda",
    AT_WRAP([&]() {
      kernel_impl<scalar_t>(iter);
    }),
    AT_EXPAND(AT_ALL_TYPES),
    AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES),
    kBFloat16, kHalf, kBool
);

这种写法在仓库真实内核中已有先例,例如 CPU 端 fill 内核 aten/src/ATen/native/cpu/FillKernel.cpp 的分发就同时列出了 AT_EXPAND(AT_ALL_TYPES_AND_COMPLEX)kBoolAT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES),说明"全类型组 + 无符号组 + 零散单类型"的混排是被生产代码接受的规范形态。除 FillKernel 外,CopyKernel.cppIndexKernel.cppScatterGatherKernel.cppBinaryOpsKernel.cppSortingKernel.cpp 以及 CUDA 侧的 CompareEQKernel.cuIndexKernel.cuShape.cu 等 20 多个内核文件(见 aten/src/ATen/native)都已在用同样的方式覆盖无符号类型,可以打开这些文件对照自己的改法。

方法二:用 AT_INTEGRAL_TYPES_V2 替换 AT_INTEGRAL_TYPES

仅适用于当前分发明确使用了 AT_EXPAND(AT_INTEGRAL_TYPES) 的场景——用其超集替换,一行搞定,更简洁:

// Before
AT_DISPATCH_V2(
    dtype,
    "integral_op",
    AT_WRAP([&]() {
      kernel<scalar_t>();
    }),
    AT_EXPAND(AT_INTEGRAL_TYPES)
);

// After(替换为 V2 超集)
AT_DISPATCH_V2(
    dtype,
    "integral_op",
    AT_WRAP([&]() {
      kernel<scalar_t>();
    }),
    AT_EXPAND(AT_INTEGRAL_TYPES_V2)
);

从定义看,AT_INTEGRAL_TYPES_V2 展开后恰好等于 AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)torch/headeronly/core/Dispatch_v2.h),所以方法二是方法一的语法糖,两者最终生成的 switch 分支完全一致。

两种方法的选择标准

条件 推荐改法
分发中出现了 AT_EXPAND(AT_INTEGRAL_TYPES) 方法二:原地替换为 AT_INTEGRAL_TYPES_V2,更简洁
分发用的是 AT_ALL_TYPES 或其他组合 方法一:追加 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)
分发现已含 AT_INTEGRAL_TYPES_V2AT_BAREBONES_UNSIGNED_TYPES 已有 uint 支持,直接跳过,不要重复添加

可以归纳成一棵决策树:

Is the file using AT_DISPATCH_V2?
├─ No → 先转换为 V2 形态,再继续
└─ Yes
   └─ 是否使用了 AT_EXPAND(AT_INTEGRAL_TYPES)?
      ├─ 是 → 替换为 AT_EXPAND(AT_INTEGRAL_TYPES_V2)
      └─ 否 → 在类型列表中追加 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)

常见模式的改前/改后对照

模式一:AT_ALL_TYPES + 零散类型。零散单类型(kHalfkBFloat16kBool 等)保持原样,无符号组插入其中即可:

// Before
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);

// After
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES), kHalf, kBFloat16);

模式二:INTEGRAL 与 FLOATING 分开列举。此时优先用方法二,只动整数部分,浮点部分不受影响:

// Before
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_FLOATING_TYPES));

// After(优先方法二:整数组升级为 V2)
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_INTEGRAL_TYPES_V2), AT_EXPAND(AT_FLOATING_TYPES));

模式三:先转换再扩展。即上文"旧宏 → V2 → 加 uint"的两步流程,适用于尚未迁移到 AT_DISPATCH_V2 的存量文件。

多分发点:一致性是正确性的关键

同一个算子往往有多个分发点——CPU 与 CUDA 各一份、同一文件里 min 和 max 两个函数、iter.dtype()iter.input_dtype() 不同入口等。必须检查文件内的全部分发点并做相同的类型覆盖更新,遗漏任一处都会导致"同一算子在不同路径上 dtype 支持不一致":

void min_values_kernel_cuda(TensorIterator& iter) {
  AT_DISPATCH_V2(iter.dtype(), "min_values_cuda", AT_WRAP([&]() {
    impl<scalar_t>(iter);
  }), AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES), kBFloat16, kHalf);
  //                           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  //                           新增的 uint 支持
}

void min_launch_kernel(TensorIterator &iter) {
  AT_DISPATCH_V2(iter.input_dtype(), "min_cuda", AT_WRAP([&]() {
    gpu_reduce_kernel<scalar_t>(iter);
  }), AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES), kBFloat16, kHalf);
  //                           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  //                           这里同样补充了 uint 支持
}

实际操作中建议直接在目标文件内全文搜索 AT_DISPATCH_V2(AT_DISPATCH_ALL_TYPES 等关键词,枚举出所有分发点逐一核对。

边界情况

情况一:纯浮点算子不要加

若算子在语义上只支持浮点(如归一化、激活函数),保持原样,不要为了"覆盖更全"而塞入无符号类型:

// 保持不动——纯浮点算子
AT_DISPATCH_V2(dtype, "float_op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf);

情况二:与复数类型共存

无符号类型与复数类型可以并列出现在同一个分发列表中,互不冲突:

AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES),
    AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES),
    AT_EXPAND(AT_COMPLEX_TYPES),
    kHalf, kBFloat16);

情况三:已经支持 uint 则跳过

改动前先检查:分发列表中已出现 AT_INTEGRAL_TYPES_V2AT_BAREBONES_UNSIGNED_TYPES,说明该算子已覆盖无符号类型,重复添加虽不致命(会产生重复 case 分支)但属于冗余改动,应直接跳过。

验证清单与测试预期

改动完成后,逐项核对以下清单:

  • [ ] 分发宏为 AT_DISPATCH_V2 形态(而非旧版 AT_DISPATCH_...);
  • [ ] 无符号类型通过上述两种方法之一加入;
  • [ ] 文件内所有相关分发点均已更新,类型覆盖一致;
  • [ ] 类型组一律用 AT_EXPAND() 包裹;
  • [ ] 各实参之间逗号分隔正确,lambda 体已被 AT_WRAP 包住。

功能层面的验证方式是构造 torch.uint16torch.uint32torch.uint64 张量后直接调用该算子:改前会因 dtype 未覆盖而在分发默认分支报错,改后应正常返回结果。仓库中与无符号类型相关的内核行为可进一步在 test/test_torch.pytest/test_type_promotion.py 等测试文件所在的 test 目录下按算子名定位;功能正确性验证由使用者按自身算子的语义自行设计用例。

小结

为 PyTorch 算子扩展 uint16/32/64 支持本质上是一类机械性强但有明确判据的分发宏改写:

  1. 确认 AT_DISPATCH_V2 前提,必要时先做 V2 化转换;
  2. 看清当前类型组——AT_ALL_TYPES 并不包含无符号 16/32/64;
  3. AT_EXPAND(AT_INTEGRAL_TYPES) 就替换成 AT_EXPAND(AT_INTEGRAL_TYPES_V2),否则追加 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)
  4. 覆盖文件中全部分发点,保持类型覆盖一致;
  5. 用三种无符号 dtype 的张量实际调用算子验证。

需要记住的两个高频误区:一是误以为 AT_ALL_TYPES "已经全了",漏掉无符号组;二是把 8 位无符号 kByte 当成 barebones 无符号类型——它本就属于 AT_INTEGRAL_TYPES,不需要也不应该重复添加。相关宏的权威定义可回到 aten/src/ATen/Dispatch_v2.htorch/headeronly/core/Dispatch_v2.h 两个文件核对,旧版宏族的展开逻辑在 aten/src/ATen/Dispatch.h

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