add-uint-support

add-uint-support

热门

通过更新 AT_DISPATCH 宏,为 PyTorch 算子添加无符号整数(uint)类型支持。适用于为算子或 Kernel 补充 uint16、uint32、uint64 类型支持的场景,或者当用户提到开启无符号类型支持、基础无符号类型(barebones unsigned types)或 uint 支持时使用。

10万Star
2.9万Fork
更新于 2026/8/4
SKILL.md
只读
名称
add-uint-support
描述

通过更新 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, kUInt64
  • AT_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_TYPES
  • AT_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 支持的任务时:

  1. 读取目标文件
  2. 检查是否已使用 AT_DISPATCH_V2:
    • 若未改用,先调用 at-dispatch-v2 Skill 转换
  3. 识别文件中所有的分派宏位置
  4. 针对每个分派宏:
    • 分析现有的类型组
    • 选择处理方式(追加 BAREBONES_UNSIGNED 或升级为 V2)
    • 使用编辑工具应用修改
  5. 向用户展示改写后的代码差异
  6. 解释具体的修改内容

注意事项

  • 务必优先确认是否需要先迁移到 V2
  • 确保全文件内所有分派点的修改方式保持一致
  • 条件允许时,方式 2(使用 AT_INTEGRAL_TYPES_V2)代码更精简
  • 方式 1(显式写出 AT_BAREBONES_UNSIGNED_TYPES)意图更明确
  • 注意:此处指的无符号类型为 kUInt16、kUInt32、kUInt64(不含代表 uint8 的 kByte)
  • 部分算子在语义上可能并不适合无符号类型,请结合实际语意做出判断

测试说明

完成 uint 支持添加后,算子应能正常接收 uint16、uint32 和 uint64 张量。具体的功能性测试由用户自行负责。