add-uint-support

add-uint-support

熱門

透過更新 AT_DISPATCH 巨集,為 PyTorch 運算子新增無符號整數(uint)型別支援。當需要為運算子、Kernel 新增 uint16、uint32、uint64 型別支援,或是使用者提及啟用無符號型別、基礎無符號型別(barebones unsigned types)或 uint 支援時使用。

10萬星標
2.9萬分支
更新於 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 支援
  • 使用者提到「無符號型別(unsigned types)」、「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

操作指南

步驟 1:判斷是否需要轉換為 V2

檢查檔案是否已使用 AT_DISPATCH_V2:

若仍使用舊版 AT_DISPATCH:

  • 先使用 at-dispatch-v2 Skill 轉換為 AT_DISPATCH_V2
  • 接著再繼續進行新增 uint 支援的操作

若已使用 AT_DISPATCH_V2:

  • 直接跳至步驟 2

步驟 2:分析目前的 dispatch 巨集

確認目前使用了哪些型別群組:

AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
  // 程式體
}), 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) → 浮點數型別

步驟 3:選擇新增 uint 的方法

有兩種做法:

方法 1:顯式新增 AT_BAREBONES_UNSIGNED_TYPES

  • 使用時機:希望明確展示已新增 uint 支援
  • 在型別清單中加入 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)

方法 2:將 AT_INTEGRAL_TYPES 替換為 AT_INTEGRAL_TYPES_V2

  • 使用時機:dispatch 中已包含 AT_EXPAND(AT_INTEGRAL_TYPES)
  • 更簡潔:直接將該型別群組替換為其超集(superset)
  • 僅適用於原本就存在 AT_INTEGRAL_TYPES 的情況

步驟 4:套用程式碼轉換

方法 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)
);

步驟 5:處理 AT_ALL_TYPES 與個別型別群組

若 dispatch 使用了 AT_EXPAND(AT_ALL_TYPES)

  • AT_ALL_TYPES = AT_INTEGRAL_TYPES + AT_FLOATING_TYPES
  • 若要新增 uint:需在清單中加上 AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES)

若 dispatch 分開列出了 INTEGRAL 與 FLOATING:

// 變更前
AT_EXPAND(AT_INTEGRAL_TYPES), AT_EXPAND(AT_FLOATING_TYPES)

// 變更後(優先使用方法 2)
AT_EXPAND(AT_INTEGRAL_TYPES_V2), AT_EXPAND(AT_FLOATING_TYPES)

步驟 6:檢查所有 dispatch 呼叫點

檢查檔案中所有需要 uint 支援的 dispatch 巨集:

  • 部分運算子可能有多個 dispatch 點(如 CPU、CUDA 或不同的功能函式)
  • 確保所有位址都一致套用此轉換
  • 確認每個呼叫點都有更新至相同的型別涵蓋範圍

步驟 7:驗證變更內容

檢查以下項目:

  • [ ] 是否已使用 AT_DISPATCH_V2 格式(非舊版 AT_DISPATCH)
  • [ ] 是否已透過兩種方法之一新增無符號型別
  • [ ] 檔案中所有相關的 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:獨立的 INTEGRAL + FLOATING

// 變更前
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:舊版 Dispatch 需要先進行轉換

// 變更前(需先轉換至 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);

多重 Dispatch 呼叫點範例

針對包含多個函式的檔案:

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:僅支援浮點數型別的 Dispatch

若該運算子僅支援浮點數型別,切勿新增 uint 支援:

// 保持原樣 - 僅限浮點數運算子
AT_DISPATCH_V2(dtype, "float_op", AT_WRAP([&]() {
  kernel<scalar_t>();
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf);

情況 2:存在複數型別

無符號型別可與複數型別共存:

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. 找出所有 dispatch 巨集呼叫點
  4. 針對每個 dispatch 進行:
    • 分析目前的型別群組
    • 選擇方法(新增 BAREBONES_UNSIGNED 或升級至 V2)
    • 使用 Edit 工具套用變更
  5. 向使用者展示變更
  6. 說明修改內容

重要注意事項

  • 務必先檢查是否需要轉換為 v2
  • 確保檔案中所有 dispatch 呼叫點均套用一致的變更
  • 在適用的情況下,方法 2(AT_INTEGRAL_TYPES_V2)會更加乾淨
  • 方法 1(顯式使用 AT_BAREBONES_UNSIGNED_TYPES)則更為直觀明確
  • 無符號型別包含:kUInt16, kUInt32, kUInt64(不含 kByte,因為 kByte 為 uint8)
  • 部分運算子可能在語意上不支援無符號型別 - 請依專業經驗判斷

測試

新增 uint 支援後,運算子應能接受 uint16、uint32 與 uint64 的 Tensor。功能測試由使用者自行負責。