at-dispatch-v2

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)實作。

10萬星標
2.9萬分支
更新於 2026/8/4
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);

核心重構重點

  1. 調整參數順序:先放置 scalar_typename,接著放置 lambda,最後才是型別(types)
  2. 包覆 lambda:使用 AT_WRAP(lambda) 來避免內部逗號引發解析問題
  3. 展開型別群組:使用 AT_EXPAND(AT_ALL_TYPES) 替代原先的隱式展開
  4. 列出個別型別:在展開的型別群組之後,依序追加額外的型別(如 kHalf、kBFloat16 等)
  5. 加入標頭檔:在其他 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 的優點

  1. 巨集名稱不再包含引數數量(arity):無需針對 AND2、AND3、AND4 使用不同的巨集
  2. 可組合的型別集合:可透過 AT_EXPAND() 自由搭配與組合各種型別群組
  3. 擴充性更佳:能輕鬆擴充更多型別,不受巨集引數限制
  4. 語意更清晰:型別群組為明確宣告,而非隱式藏於巨集名稱中

重要注意事項

  • 請保留 #include <ATen/Dispatch.h>——其他程式碼可能仍需使用
  • AT_WRAP() 為強制使用——可防止 lambda 內部的逗號導致解析錯誤
  • 型別群組必須加上 AT_EXPAND(),個別型別則不需要
  • v2 API 定義於 aten/src/ATen/Dispatch_v2.h——完整說明請參閱該檔案
  • 請參閱該標頭檔以取得用於重新產生巨集實作的 Python 指令碼

工作流程

收到轉換 AT_DISPATCH 巨集的要求時:

  1. 讀取檔案並找出所有使用 AT_DISPATCH 的地方
  2. 若尚未引入,請新增 #include <ATen/Dispatch_v2.h>
  3. 針對每個 dispatch 巨集:
    • 辨識其模式並提取構成元件
    • 對映基底型別群組
    • 擷取個別型別
    • 構建 AT_DISPATCH_V2 呼叫
    • 使用 Edit 工具套用修改
  4. 向使用者展示轉換完成的完整檔案
  5. 說明變更內容

請勿編譯或測試程式碼——專注於精確轉換即可。