透過更新 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, 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
操作指南
步驟 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_TYPESAT_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 支援的需求時:
- 讀取目標檔案
- 檢查是否使用 AT_DISPATCH_V2:
- 若否 → 先使用 at-dispatch-v2 Skill
- 找出所有 dispatch 巨集呼叫點
- 針對每個 dispatch 進行:
- 分析目前的型別群組
- 選擇方法(新增 BAREBONES_UNSIGNED 或升級至 V2)
- 使用 Edit 工具套用變更
- 向使用者展示變更
- 說明修改內容
重要注意事項
- 務必先檢查是否需要轉換為 v2
- 確保檔案中所有 dispatch 呼叫點均套用一致的變更
- 在適用的情況下,方法 2(AT_INTEGRAL_TYPES_V2)會更加乾淨
- 方法 1(顯式使用 AT_BAREBONES_UNSIGNED_TYPES)則更為直觀明確
- 無符號型別包含:kUInt16, kUInt32, kUInt64(不含 kByte,因為 kByte 為 uint8)
- 部分運算子可能在語意上不支援無符號型別 - 請依專業經驗判斷
測試
新增 uint 支援後,運算子應能接受 uint16、uint32 與 uint64 的 Tensor。功能測試由使用者自行負責。






