SKILL.md
唯讀
名稱
at-dispatch-v2
描述
將 ATen C++ 程式碼中的 PyTorch AT_DISPATCH 巨集轉換為 AT_DISPATCH_V2 格式。適用於將 AT_DISPATCH_ALL_TYPES_AND*、AT_DISPATCH_FLOATING_TYPES* 或其他分發巨集(dispatch macros)移植至全新 v2 API 的情境。主要用於 ATen kernel 檔案、CUDA kernel 與原生算子(native operator)實作。
AT_DISPATCH 轉 AT_DISPATCH_V2 轉換器
本 Skill 旨在協助將 PyTorch 舊有的 AT_DISPATCH 巨集轉換為 aten/src/ATen/Dispatch_v2.h 中定義的新版 AT_DISPATCH_V2 格式。
何時使用此 Skill
在以下情境下使用此 Skill:
- 需要將 AT_DISPATCH_* 巨集轉換為 AT_DISPATCH_V2
- 正在移植 ATen kernel 以使用新版 dispatch API
- 處理
aten/src/ATen/native/目錄中包含 dispatch 巨集的檔案 - 使用者提到 "AT_DISPATCH"、"dispatch v2"、"Dispatch_v2.h" 或巨集轉換
快速參考
舊格式:
AT_DISPATCH_ALL_TYPES_AND3(kBFloat16, kHalf, kBool, dtype, "kernel_name", [&]() {
// lambda body
});
新格式:
AT_DISPATCH_V2(dtype, "kernel_name", AT_WRAP([&]() {
// lambda body
}), AT_EXPAND(AT_ALL_TYPES), kBFloat16, kHalf, kBool);
核心重構重點
- 調整參數順序:先放置
scalar_type與name,接著放置 lambda,最後才是型別(types) - 包覆 lambda:使用
AT_WRAP(lambda)來避免內部逗號引發解析問題 - 展開型別群組:使用
AT_EXPAND(AT_ALL_TYPES)替代原先的隱式展開 - 列出個別型別:在展開的型別群組之後,依序追加額外的型別(如 kHalf、kBFloat16 等)
- 加入標頭檔:在其他 Dispatch 標頭檔旁加上
#include <ATen/Dispatch_v2.h>
操作步驟
步驟 1:加入 Dispatch_v2.h 標頭檔
在現有的 #include <ATen/Dispatch.h> 旁加上 v2 標頭檔:
#include <ATen/Dispatch.h>
#include <ATen/Dispatch_v2.h>
暫時保留舊有的 Dispatch.h(其他程式碼可能仍會用到)。
步驟 2:辨識舊版 dispatch 模式
常見需要轉換的模式:
AT_DISPATCH_ALL_TYPES_AND{2,3,4}(type1, type2, ..., scalar_type, name, lambda)AT_DISPATCH_FLOATING_TYPES_AND{2,3}(type1, type2, ..., scalar_type, name, lambda)AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND{2,3}(type1, ..., scalar_type, name, lambda)AT_DISPATCH_FLOATING_AND_COMPLEX_TYPES_AND{2,3}(type1, ..., scalar_type, name, lambda)
步驟 3:將舊版巨集對映至型別群組
對照舊版基底型別巨集對應的型別群組:
| 舊版巨集基底 | AT_DISPATCH_V2 型別群組 |
|---|---|
ALL_TYPES |
AT_EXPAND(AT_ALL_TYPES) |
FLOATING_TYPES |
AT_EXPAND(AT_FLOATING_TYPES) |
INTEGRAL_TYPES |
AT_EXPAND(AT_INTEGRAL_TYPES) |
COMPLEX_TYPES |
AT_EXPAND(AT_COMPLEX_TYPES) |
ALL_TYPES_AND_COMPLEX |
AT_EXPAND(AT_ALL_TYPES_AND_COMPLEX) |
對於組合模式,請使用多個 AT_EXPAND() 條目:
// 舊版:AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2(...)
// 新版:AT_EXPAND(AT_ALL_TYPES), AT_EXPAND(AT_COMPLEX_TYPES), type1, type2
步驟 4:擷取個別型別
從 AT_DISPATCH_*_AND2(type1, type2, ...) 或 AT_DISPATCH_*_AND3(type1, type2, type3, ...) 中提取出個別的型別(type1、type2 等)。
這些型別將作為接在型別群組後方的末尾參數:
AT_DISPATCH_V2(..., AT_EXPAND(AT_ALL_TYPES), kBFloat16, kHalf, kBool)
^^^^^^^^^^^^^^^^^^^^^^^^
來自 AND3 的個別型別
步驟 5:轉換為 AT_DISPATCH_V2
執行以下轉換:
模式:
AT_DISPATCH_V2(
scalar_type, // 第 1 個:dtype 運算式
"name", // 第 2 個:除錯名稱字串
AT_WRAP(lambda), // 第 3 個:經由 AT_WRAP 包覆的 lambda
type_groups, // 第 4 個以上:使用 AT_EXPAND() 的型別群組
individual_types // 最後:個別型別
)
轉換範例:
// 轉換前
AT_DISPATCH_ALL_TYPES_AND3(
kBFloat16, kHalf, kBool,
iter.dtype(),
"min_values_cuda",
[&]() {
min_values_kernel_cuda_impl<scalar_t>(iter);
}
);
// 轉換後
AT_DISPATCH_V2(
iter.dtype(),
"min_values_cuda",
AT_WRAP([&]() {
min_values_kernel_cuda_impl<scalar_t>(iter);
}),
AT_EXPAND(AT_ALL_TYPES),
kBFloat16, kHalf, kBool
);
步驟 6:處理多行 Lambda
若 lambda 內部包含逗號或複雜運算式,AT_WRAP 至關重要:
AT_DISPATCH_V2(
dtype,
"complex_kernel",
AT_WRAP([&]() {
gpu_reduce_kernel<scalar_t, scalar_t>(
iter,
MinOps<scalar_t>{},
thrust::pair<scalar_t, int64_t>(upper_bound(), 0) // 內部含有逗號!
);
}),
AT_EXPAND(AT_ALL_TYPES)
);
步驟 7:驗證轉換結果
檢查以下項目:
- [ ]
AT_WRAP()已包覆整個 lambda - [ ] 型別群組均使用了
AT_EXPAND() - [ ] 個別型別沒有加上
AT_EXPAND()(僅保持kBFloat16,而非AT_EXPAND(kBFloat16)) - [ ] 參數順序為:scalar_type、name、lambda、types
- [ ] 已加上標頭檔:
#include <ATen/Dispatch_v2.h>
型別群組參考清單
可用的型別群組巨集(需搭配 AT_EXPAND() 使用):
AT_INTEGRAL_TYPES // kByte, kChar, kInt, kLong, kShort
AT_FLOATING_TYPES // kDouble, kFloat
AT_COMPLEX_TYPES // kComplexDouble, kComplexFloat
AT_QINT_TYPES // kQInt8, kQUInt8, kQInt32
AT_ALL_TYPES // INTEGRAL_TYPES + FLOATING_TYPES
AT_ALL_TYPES_AND_COMPLEX // ALL_TYPES + COMPLEX_TYPES
AT_INTEGRAL_TYPES_V2 // INTEGRAL_TYPES + 無號數型別 (unsigned types)
AT_BAREBONES_UNSIGNED_TYPES // kUInt16, kUInt32, kUInt64
AT_FLOAT8_TYPES // Float8 變體
常見轉換模式
模式:AT_DISPATCH_ALL_TYPES_AND2
// 轉換前
AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBFloat16, dtype, "op", [&]() {
kernel<scalar_t>(data);
});
// 轉換後
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>(data);
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBFloat16);
模式:AT_DISPATCH_FLOATING_TYPES_AND3
// 轉換前
AT_DISPATCH_FLOATING_TYPES_AND3(kHalf, kBFloat16, kFloat8_e4m3fn,
tensor.scalar_type(), "float_op", [&] {
process<scalar_t>(tensor);
});
// 轉換後
AT_DISPATCH_V2(tensor.scalar_type(), "float_op", AT_WRAP([&] {
process<scalar_t>(tensor);
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf, kBFloat16, kFloat8_e4m3fn);
模式:AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2
// 轉換前
AT_DISPATCH_ALL_TYPES_AND_COMPLEX_AND2(
kComplexHalf, kHalf,
self.scalar_type(),
"complex_op",
[&] {
result = compute<scalar_t>(self);
}
);
// 轉換後
AT_DISPATCH_V2(
self.scalar_type(),
"complex_op",
AT_WRAP([&] {
result = compute<scalar_t>(self);
}),
AT_EXPAND(AT_ALL_TYPES),
AT_EXPAND(AT_COMPLEX_TYPES),
kComplexHalf,
kHalf
);
邊界情況(Edge Cases)
情況 1:無額外型別(較罕見)
// 轉換前
AT_DISPATCH_ALL_TYPES(dtype, "op", [&]() { kernel<scalar_t>(); });
// 轉換後
AT_DISPATCH_V2(dtype, "op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES));
情況 2:包含多個個別型別(AND4、AND5 等)
// 轉換前
AT_DISPATCH_FLOATING_TYPES_AND4(kHalf, kBFloat16, kFloat8_e4m3fn, kFloat8_e5m2,
dtype, "float8_op", [&]() { kernel<scalar_t>(); });
// 轉換後
AT_DISPATCH_V2(dtype, "float8_op", AT_WRAP([&]() {
kernel<scalar_t>();
}), AT_EXPAND(AT_FLOATING_TYPES), kHalf, kBFloat16, kFloat8_e4m3fn, kFloat8_e5m2);
情況 3:不進行擷取的 Lambda
// 轉換前
AT_DISPATCH_ALL_TYPES_AND2(kHalf, kBool, dtype, "op", []() {
static_kernel<scalar_t>();
});
// 轉換後
AT_DISPATCH_V2(dtype, "op", AT_WRAP([]() {
static_kernel<scalar_t>();
}), AT_EXPAND(AT_ALL_TYPES), kHalf, kBool);
AT_DISPATCH_V2 的優點
- 巨集名稱不再包含引數數量(arity):無需針對 AND2、AND3、AND4 使用不同的巨集
- 可組合的型別集合:可透過
AT_EXPAND()自由搭配與組合各種型別群組 - 擴充性更佳:能輕鬆擴充更多型別,不受巨集引數限制
- 語意更清晰:型別群組為明確宣告,而非隱式藏於巨集名稱中
重要注意事項
- 請保留
#include <ATen/Dispatch.h>——其他程式碼可能仍需使用 AT_WRAP()為強制使用——可防止 lambda 內部的逗號導致解析錯誤- 型別群組必須加上
AT_EXPAND(),個別型別則不需要 - v2 API 定義於
aten/src/ATen/Dispatch_v2.h——完整說明請參閱該檔案 - 請參閱該標頭檔以取得用於重新產生巨集實作的 Python 指令碼
工作流程
收到轉換 AT_DISPATCH 巨集的要求時:
- 讀取檔案並找出所有使用
AT_DISPATCH的地方 - 若尚未引入,請新增
#include <ATen/Dispatch_v2.h> - 針對每個 dispatch 巨集:
- 辨識其模式並提取構成元件
- 對映基底型別群組
- 擷取個別型別
- 構建
AT_DISPATCH_V2呼叫 - 使用 Edit 工具套用修改
- 向使用者展示轉換完成的完整檔案
- 說明變更內容
請勿編譯或測試程式碼——專注於精確轉換即可。






