使用 TileLang 設計高效能 GPU 核心:Tensor-Core GEMM、融合 Softmax、FlashAttention 與自動調校
重點速覽
- •TileLang 讓開發者能在 Python 中以 tile 層級撰寫最佳化 GPU 核心,同時由編譯器自動處理執行緒映射、記憶體配置、同步與 CUDA 指令生成。
- •該教學逐步實作並驗證了向量加法、tensor-core 矩陣乘法、帶 bias 與 GELU 的融合 GEMM epilogue、逐列 softmax,以及 FlashAttention 等核心,並與 PyTorch 基準進行比較。
- •TileLang 的自動調校 decorator 會在 tile 大小、管線深度與執行緒數量之間搜尋,以自動找出依架構而定的最佳配置,解決 Ampere 與 Hopper 等不同 GPU 世代之間理想參數會變動的挑戰。
- •教學展示的核心融合技術可在寫入最終輸出前,於常駐暫存器的累加器中完成 bias 加法與 GELU 啟用等運算,從而降低中間全域記憶體流量。
- •FlashAttention 實作在不具體化完整 attention-score 矩陣的情況下處理 query、key 與 value tile,並使用線上 softmax 更新將記憶體佔用從平方空間降低到線性空間。
- •生成原始碼檢查、device-side printing 與內建 profiler 共同讓開發者能觀察編譯器輸出的 tensor-core operations、asynchronous copies 與 synchronization barriers,以支援除錯與最佳化。

TileLang 是一種高階 Python 領域特定語言(DSL),用於透過 TVM 設計並編譯以效能為導向的 GPU 核心。隨著大型語言模型與其他 transformer 架構推高運算需求,撰寫能充分利用 tensor cores(NVIDIA 專用矩陣乘法加速器)的自訂 GPU 核心,對訓練與推論效率而言變得日益關鍵。然而,傳統上撰寫這類核心需要深厚的 CUDA C++、warp 層級程式設計,以及硬體特定記憶體階層方面的專業知識。TileLang 透過讓開發者在 tile 層級表達運算,並由編譯器處理底層細節,來填補這項落差。本教學完整介紹 TileLang 的能力,從環境驗證以及建立可重複使用的基準測試與數值驗證工具開始。接著逐步實作向量加法、分塊 tensor-core 矩陣乘法、排程探索、融合 GEMM epilogue、逐列 softmax,以及 FlashAttention。
在整個教學中,開發者會直接使用 TileLang 的共享記憶體 tile、暫存器片段、管線化迴圈、平行迭代原語、歸約,以及 tensor-core GEMM 運算子。編譯器會自動處理執行緒映射、記憶體配置、同步、向量化與底層 CUDA 指令生成。核心會與 PyTorch 和 cuBLAS 基準進行比較,並檢查生成的 CUDA 原始碼、評估記憶體與運算吞吐量,以及使用自動調校來找出依架構而定的核心配置。TileLang 可在 GitHub 取得。
教學首先設定 Google Colab CUDA 環境,安裝 TileLang 並提供 nightly fallback,然後匯入所需的 PyTorch 與 TileLang 模組。接著定義可重複使用的基準測試、驗證與報告工具,用於測量核心延遲,並以相對誤差比較數值輸出。隨後實作一個 TileLang 向量加法核心,在 GPU 上執行,與 PyTorch 的頻寬進行比較,並檢查編譯器生成的 CUDA 原始碼。
接下來,教學實作一個分塊 tensor-core 矩陣乘法核心,將輸入 tile 在全域記憶體、共享記憶體與暫存器片段之間移動。開發者可手動控制 tile 維度、管線階段、執行緒數量與 L2 swizzling,同時由 TileLang 生成 tensor-core 指令、同步與記憶體傳輸邏輯。教學對多種排程配置進行基準測試,驗證其數值準確性,並找出效能最高且依架構而定的核心配置。完整教學程式碼可在此處取得。
矩陣乘法核心隨後被擴充,將 bias 加法與 GELU 啟用函式直接融合到常駐於暫存器的累加器中。此方法透過在寫入最終輸出張量之前完成 epilogue,降低中間全域記憶體流量。這類核心融合是降低記憶體頻寬瓶頸的成熟技術,而記憶體頻寬瓶頸通常主導大規模神經網路工作負載的延遲。融合實作會與 eager PyTorch 執行進行比較。此外,教學也使用片段層級的最大值與總和歸約實作逐列 softmax 核心,使正規化流程大多維持在暫存器內完成。
教學實作一個融合的 FlashAttention forward 核心,用於處理 query、key 與 value tile,而不在全域記憶體中具體化完整的 attention-score 矩陣。線上 softmax 更新透過 running maxima、normalization sums、rescaling factors,以及分塊 tensor-core 矩陣乘法來套用。FlashAttention 最初由 Stanford 研究人員提出,已成為一種廣泛採用的技術,可將 self-attention 的記憶體佔用從平方空間降低到線性空間,並且目前已整合到包括 PyTorch scaled dot-product attention API 在內的主要框架中。因果與非因果 attention 都會與 PyTorch scaled dot-product attention 進行驗證,並比較其延遲與運算吞吐量。
教學定義了一個自動調校搜尋空間,涵蓋矩陣 tile 大小、K-block 維度、管線深度與執行緒數量,同時過濾掉超出共享記憶體預算的配置。TileLang 的自動調校 decorator 用於為同一個矩陣乘法工作負載編譯、基準測試、驗證並快取多個核心排程。選定的核心會被執行,與 PyTorch 進行驗證,並評估其達成的延遲與 tensor-core 吞吐量。這項自動化搜尋解決了 GPU 核心開發中的實務挑戰:最佳 tile 大小與管線深度會因 Ampere、Hopper 等 GPU 架構而異,當目標涵蓋多個硬體世代時,手動調校容易變得脆弱。
教學透過 device-side printing、生成 CUDA 檢查,以及內建核心 profiler 介紹 TileLang 的除錯與內省工作流程。文中檢視編譯器輸出的關鍵標記,包括 tensor-core operations、asynchronous copies、synchronization barriers,以及 matrix-load instructions。所有教學段落都被組織到一個具容錯能力的 runner 中,用於記錄執行狀態、回報計時資訊,並列印精簡的 TileLang 程式設計參考。
本教學展示了 TileLang 如何將 tile 層級的 Python 程式轉換為最佳化 GPU 核心,而不需要手動管理執行緒索引、warp 層級資料配置、tensor-core 指令或非同步記憶體屏障。已實作並驗證的核心涵蓋頻寬受限的 elementwise 運算、運算密集型 GEMM 工作負載、融合神經網路 epilogue、常駐暫存器的歸約,以及 online-softmax attention。這項探索也說明了區塊維度、共享記憶體消耗、管線深度、執行緒數量、tile 形狀與 L2 swizzling 如何影響不同 GPU 架構上的效能。生成原始碼檢查、device-side 除錯、profiling 與自動化排程搜尋,共同建立了一套用於開發、驗證、基準測試與改進自訂 TileLang 核心的完整工作流程。