通过更新 AT_DISPATCH 宏,为 PyTorch 算子添加无符号整数(uint)类型支持。适用于为算子或 Kernel 补充 uint16、uint32、uint64 类型支持的场景,或者当用户提到开启无符号类型支持、基础无符号类型(barebones unsigned types)或 uint 支持时使用。
为算子添加无符号整数(uint)类型支持
本 Skill 用于指导通过更新 AT_DISPATCH 宏,为 PyTorch 算子补充无符号整数类型(uint16、uint32、uint64)支持。
适用场景
当出现以下情况时使用此 Skill:
- 需要为某个算子添加 uint16、uint32 或 uint64 支持
- 用户提到了“无符号类型”、“uint 支持”或“基础无符号类型(barebones unsigned types)”
- 需要在 Kernel 中开启对 kUInt16、kUInt32、kUInt64 的支持
- 正在处理需要扩展类型覆盖范围的算子实现
快速参考
给现有的分派(Dispatch)添加无符号类型:
// 修改前
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES));
// 修改后(方式 1:显式追加无符号类型)
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES));
// 修改后(方式 2:若存在 AT_INTEGRAL_TYPES,改用 V2 整型组)
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_INTEGRAL_TYPES_V2), AT_EXPAND(AT_FLOATING_TYPES));
类型组参考
无符号类型组:
AT_BAREBONES_UNSIGNED_TYPES: kUInt16, kUInt32, kUInt64AT_INTEGRAL_TYPES_V2: AT_INTEGRAL_TYPES + AT_BAREBONES_UNSIGNED_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
操作指南
第一步:判断是否需要先改用 V2
检查文件当前是否使用的是 AT_DISPATCH_V2:
如果仍在使用旧版 AT_DISPATCH:
- 先调用 at-dispatch-v2 Skill 将其转换为 AT_DISPATCH_V2
- 转换完成后再继续添加 uint 支持
如果已经使用了 AT_DISPATCH_V2:
- 直接进入第二步
第二步:分析当前的分派宏
明确目前已支持的类型组:
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
// body
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);
^^^^^^^^^^^^^^^^^^^^^^^^^
当前覆盖的类型
常见模式:
AT_EXPAND(AT_ALL_TYPES)→ 包含 AT_INTEGRAL_TYPES + AT_FLOATING_TYPESAT_EXPAND(AT_INTEGRAL_TYPES)→ 仅包含有符号整数AT_EXPAND(AT_FLOATING_TYPES)→ 浮点类型
第三步:选择添加 uint 支持的方式
两种实现方式:
方式 1:显式追加 AT_BAREBONES_UNSIGNED_TYPES
- 适用场景:需要清晰明确地声明追加了 uint 支持
- 做法:在类型列表中追加
AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)
方式 2:将 AT_INTEGRAL_TYPES 替换为 AT_INTEGRAL_TYPES_V2
- 适用场景:当前分派宏中已经使用了
AT_EXPAND(AT_INTEGRAL_TYPES) - 更加简练:用超集类型组直接替换原类型组
- 仅适用于原宏中包含 AT_INTEGRAL_TYPES 的情况
第四步:执行代码改写
方式 1 示例:
// 修改前
AT_DISPATCH_V2(
dtype,
"min_values_cuda",
AT_WRAP([&]() {
kernel_impl<scalar_t>(iter);
}),
AT_EXPAND(AT_ALL_TYPES),
kBFloat16, kHalf, kBool
);
// 修改后(追加无符号类型)
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
);
方式 2 示例:
// 修改前
AT_DISPATCH_V2(
dtype,
"integral_op",
AT_WRAP([&]() {
kernel<scalar_t>();
}),
AT_EXPAND(AT_INTEGRAL_TYPES)
);
// 修改后(替换为 V2)
AT_DISPATCH_V2(
dtype,
"integral_op",
AT_WRAP([&]() {
kernel<scalar_t>();
}),
AT_EXPAND(AT_INTEGRAL_TYPES_V2)
);
第五步:区分处理 AT_ALL_TYPES 与独立类型组
如果分派宏使用了 AT_EXPAND(AT_ALL_TYPES):
AT_ALL_TYPES=AT_INTEGRAL_TYPES+AT_FLOATING_TYPES- 添加 uint:直接在列表中加上
AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)
如果分派宏是分别列出 INTEGRAL 和 FLOATING:
// 修改前
AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_FLOATING_TYPES)
// 修改后(优先推荐方式 2)
AT_EXPAND(AT_INTEGRAL_TYPES_V2), AT_EXPAND(AT_FLOATING_TYPES)
第六步:核对所有分派点
检查文件中所有需要添加 uint 支持的分派宏:
- 部分算子可能会存在多个分派点(如 CPU、CUDA 或不同的辅助函数)
- 保持修改的一致性,覆盖所有相关分派点
- 确保每个分派点都更新了相同的类型支持
第七步:验证改写结果
确认满足以下条件:
- [ ] 已统一使用 AT_DISPATCH_V2 格式(而非旧版 AT_DISPATCH)
- [ ] 已通过上述两种方式之一成功添加了无符号类型
- [ ] 文件中所有相关的分派点都已更新
- [ ] 类型组外层均使用了
AT_EXPAND()包裹 - [ ] 参数格式正确,用逗号正常分隔
常见模式
模式 1:AT_ALL_TYPES + 其它扩展类型
// 修改前
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);
// 修改后
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES), kHalf, kBFloat16);
模式 2:整型与浮点分开列出
// 修改前
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_FLOATING_TYPES));
// 修改后
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_INTEGRAL_TYPES_V2), AT_EXPAND(AT_FLOATING_TYPES));
模式 3:旧版分派宏需先完成迁移
// 修改前(需先迁移至 V2)
AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBFloat16, dtype, "op", [&]() {
kernel<scalar_t>();
});
// 迁移至 V2 之后
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);
// 添加 uint 支持之后
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES), kHalf, kBFloat16);
多分派点场景示例
对于包含多个实现函数的文件:
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-v2 Skill 转换,完成后再继续
└─ 是
└─ 是否使用了 AT_EXPAND(AT_INTEGRAL_TYPES)?
├─ 是 → 替换为 AT_EXPAND(AT_INTEGRAL_TYPES_V2)
└─ 否 → 在类型列表中追加 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)
边缘情况与特例
场景 1:仅支持浮点类型的分派宏
如果算子本身仅支持浮点类型,无需添加 uint 支持:
// 保持原样——纯浮点算子
AT_DISPATCH_V2(dtype, "float_op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf);
场景 2:包含复数类型(Complex types)
无符号类型可与复数类型并行共存:
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);
场景 3:已具备 uint 支持
检查代码中是否已经包含了 uint 支持:
- 若已使用
AT_INTEGRAL_TYPES_V2→ 说明已有 uint 支持 - 若列表中已有
AT_BAREBONES_UNSIGNED_TYPES→ 说明已有 uint 支持 - 如果已经具备 uint 支持,跳过该文件即可
工作流
当接到添加 uint 支持的任务时:
- 读取目标文件
- 检查是否已使用 AT_DISPATCH_V2:
- 若未改用,先调用 at-dispatch-v2 Skill 转换
- 识别文件中所有的分派宏位置
- 针对每个分派宏:
- 分析现有的类型组
- 选择处理方式(追加 BAREBONES_UNSIGNED 或升级为 V2)
- 使用编辑工具应用修改
- 向用户展示改写后的代码差异
- 解释具体的修改内容
注意事项
- 务必优先确认是否需要先迁移到 V2
- 确保全文件内所有分派点的修改方式保持一致
- 条件允许时,方式 2(使用
AT_INTEGRAL_TYPES_V2)代码更精简 - 方式 1(显式写出
AT_BAREBONES_UNSIGNED_TYPES)意图更明确 - 注意:此处指的无符号类型为 kUInt16、kUInt32、kUInt64(不含代表 uint8 的 kByte)
- 部分算子在语义上可能并不适合无符号类型,请结合实际语意做出判断
测试说明
完成 uint 支持添加后,算子应能正常接收 uint16、uint32 和 uint64 张量。具体的功能性测试由用户自行负责。






