PyTorch 算子类型扩展实战:用 AT_DISPATCH_V2 为 uint16/uint32/uint64 补齐无符号整数支持
本文为 PyTorch 开发者讲解如何为算子内核(kernel)添加无符号整数类型(kUInt16、kUInt32、kUInt64)的分发支持:核心手法是在算子实现的 AT_DISPATCH_V2 宏中引入 AT_BAREBONES_UNSIGNED_TYPES 类型组,或把 AT_INTEGRAL_TYPES 替换为其超集 AT_INTEGRAL_TYPES_V2。读完本文,你将掌握判断何时需要转换到 V2 宏、两种改法的适用条件、多分发点的一致修改策略,以及基于 aten/src/ATen/Dispatch_v2.h 与 torch/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_AP1~AT_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
这里有两个要点需要特别注意:
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)才生效。kByte(uint8)不属于"barebones"无符号组。kByte已经包含在AT_INTEGRAL_TYPES中;AT_BAREBONES_UNSIGNED_TYPES仅指kUInt16、kUInt32、kUInt64三个类型。另外 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)、kBool 与 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES),说明"全类型组 + 无符号组 + 零散单类型"的混排是被生产代码接受的规范形态。除 FillKernel 外,CopyKernel.cpp、IndexKernel.cpp、ScatterGatherKernel.cpp、BinaryOpsKernel.cpp、SortingKernel.cpp 以及 CUDA 侧的 CompareEQKernel.cu、IndexKernel.cu、Shape.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_V2 或 AT_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 + 零散类型。零散单类型(kHalf、kBFloat16、kBool 等)保持原样,无符号组插入其中即可:
// 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_V2 或 AT_BAREBONES_UNSIGNED_TYPES,说明该算子已覆盖无符号类型,重复添加虽不致命(会产生重复 case 分支)但属于冗余改动,应直接跳过。
验证清单与测试预期
改动完成后,逐项核对以下清单:
- [ ] 分发宏为
AT_DISPATCH_V2形态(而非旧版AT_DISPATCH_...); - [ ] 无符号类型通过上述两种方法之一加入;
- [ ] 文件内所有相关分发点均已更新,类型覆盖一致;
- [ ] 类型组一律用
AT_EXPAND()包裹; - [ ] 各实参之间逗号分隔正确,lambda 体已被
AT_WRAP包住。
功能层面的验证方式是构造 torch.uint16、torch.uint32、torch.uint64 张量后直接调用该算子:改前会因 dtype 未覆盖而在分发默认分支报错,改后应正常返回结果。仓库中与无符号类型相关的内核行为可进一步在 test/test_torch.py、test/test_type_promotion.py 等测试文件所在的 test 目录下按算子名定位;功能正确性验证由使用者按自身算子的语义自行设计用例。
小结
为 PyTorch 算子扩展 uint16/32/64 支持本质上是一类机械性强但有明确判据的分发宏改写:
- 确认
AT_DISPATCH_V2前提,必要时先做 V2 化转换; - 看清当前类型组——
AT_ALL_TYPES并不包含无符号 16/32/64; - 有
AT_EXPAND(AT_INTEGRAL_TYPES)就替换成AT_EXPAND(AT_INTEGRAL_TYPES_V2),否则追加AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES); - 覆盖文件中全部分发点,保持类型覆盖一致;
- 用三种无符号 dtype 的张量实际调用算子验证。
需要记住的两个高频误区:一是误以为 AT_ALL_TYPES "已经全了",漏掉无符号组;二是把 8 位无符号 kByte 当成 barebones 无符号类型——它本就属于 AT_INTEGRAL_TYPES,不需要也不应该重复添加。相关宏的权威定义可回到 aten/src/ATen/Dispatch_v2.h 与 torch/headeronly/core/Dispatch_v2.h 两个文件核对,旧版宏族的展开逻辑在 aten/src/ATen/Dispatch.h。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0623
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00