為 AMD MI300 創建自訂核心
Hugging Face 與 AMD 開發了一組針對 AMD MI300X 的開源優化核心,以使用 VLLM 提升 Llama 3.1 405B 在 FP8 下的服務效能。透過實作三個特定的自訂核心——融合殘差連接/RMS norm/FP8 轉換核心、融合 SwiGLU 激活/FP8 轉換核心以及 Skinny GEMM 核心——團隊在解碼階段實現了顯著的延遲減少(輸入大小為 1,輸出大小為 128 進行測量)。
自訂核心實作與效能提升
融合 RMS Norm 核心
RMS norm 核心透過將殘差連接、逐行 Root Mean Square (RMS) 標準化以及 FP8 量化融合為單一操作,來優化解碼塊的開頭。
技術優化:
- 向量化記憶體存取: 該核心使用 128 位寬的載入來每指令取得 8 個 FP16 元素,確保記憶體存取是合併且連續的,以最大化 warp 效率。
- 共享記憶體 (SMEM) 利用: 為避免重複存取 VRAM(全域記憶體),隱藏狀態 $x$ 的修改版本會被存放在共享記憶體中。對於 Llama 405B,維度 $d=16384$ 能夠放入每個計算單元可用的 64KB 共享記憶體內。
- 區塊級別減少: 該核心為每行分配一個執行緒區塊,並使用共享記憶體來同步執行緒以進行 RMS norm 所需的求和。
結果: 「向量化 + SMEM」實作優於標準 PyTorch 與 VLLM 的現有實作,在各種批次大小下提供顯著的加速。
融合 SwiGLU 核心
SwiGLU 核心將激活函數及其後的 FP8 量化融合為 MLP 區塊的「Gate / Up」投影。
技術優化:
- 打包指令: 該核心利用 MI300X 的打包指令來進行 FP16 加法與乘法,以提升每指令的工作量。
- 快速數學近似: 為降低延遲,該核心將標準的
exp指令替換為較快的exp2指令,方法是將輸入縮放 $\log(2)$ 倍,這導致精度損失可忽略不計。 - 打包 FP32 到 FP8 轉換: 由於 MI300X 僅支援從 FP32 轉換為 FP8,該核心利用打包轉換指令來提升效能。
結果: 自訂 SwiGLU 核心平均比 PyTorch 快超過 14 倍,且比 VLLM 核心快 27% 到 100%。
細長 GEMM 核心
標準庫中的一般矩陣乘法 (GEMM) 核心對於「細長」矩陣——即只有很少幾行的矩陣(在批次大小低時的解碼過程中常見)——通常效率不高,因為它們會導致 GPU 利用率低,原因是瓦片機會有限。
技術優化:
- Split-K 演算法: 該核心將 GEMM 沿著共享 K 軸分割為數個並行執行的子 GEMM。這樣通過分配工作負載,增加了活躍的計算單元 (CU) 數量,減少每個 CU 在 K 軸上花費的時間。
- 用於移除填充的稀疏技巧: 為避免在行數小於最小密集張量核心指令大小(例如 16)時產生浪費的填充,該核心使用 4:2 結構化稀疏指令。一個 8 行的密集矩陣會被映射為一個 16 行的稀疏矩陣,從而可以使用
16x16x64稀疏指令,其深度是最小密集指令的兩倍。 - 執行緒專業化與非同步執行: 為應對低算術強度,該核心將 warp 分為「生產者」(專門負責從 VRAM 載入資料到共享記憶體)和「消費者」(專門負責計算)。共享記憶體中的佇列會非同步協調這些 warp,確保消費者在等待緩慢的 VRAM 載入時不會閒置。
結果: Skinny GEMM 核心在低行數(M = 1, 8, 16)時相較於 PyTorch 顯示出顯著的加速,特別是對 QKV 與 Gate/Up 投影而言,但隨著批次大小增加到 32,這些加速會逐漸減弱。
實作與可用性
所有開發的核心都可在 hf-rocm-kernels GitHub 儲存庫中取得,其中包含原始碼、Python 繫結、基準測試腳本以及測試套件。這些核心設計為可獨立使用或整合到 VLLM。為了重現結果,Hugging Face 建議使用開發期間所使用的特定 ROCm 6.3.1 容器。