CODA: 透過 GEMM-Epilogue 融合優化 Transformer 訓練

Transformer 訓練系統從根本上是建立在密集線性代數之上的,然而端到端執行時間的很大一部分通常被記憶體受限(memory-bound)的算子所消耗。雖然核心矩陣乘法(GEMMs)已經過高度優化,但周邊的操作——例如歸一化(normalization)、激活(activations)、殘差更新(residual updates)和歸約(reductions)——經常需要透過全域記憶體(global memory)移動大型中間張量。這種重複的數據移動造成了關鍵的瓶頸,限制了訓練堆疊的整體效率。

為了應對這一點,研究人員引入了 CODA,這是一種 GPU kernel 抽象,旨在將這些計算表達為「GEMM-plus-epilogue」程式。透過重新思考 Transformer 區塊的執行方式,CODA 旨在消除不必要的記憶體往返,並最大化晶片上數據的利用率。

核心問題:Transformer 中的記憶體牆

在標準的 Transformer 區塊中,操作序列通常被處理為一系列獨立的框架 kernel。例如,一個 GEMM 操作會產生一個結果並將其寫入全域記憶體,隨後的一個激活或歸一化 kernel 僅為了執行簡單的算術運算,就必須將該數據讀回 GPU 的暫存器或共享記憶體,然後再次寫回。

這種模式在本質上是低效的。因為這些「epilogue」操作在移動每個位元組的數據時所執行的算術運算非常少,因此它們是記憶體受限的。隨著硬體計算能力增長的速度快於記憶體頻寬,這些非 GEMM 操作在訓練時間中佔據了不成比例的較大比例。

CODA 的方法:GEMM-Epilogue 程式

CODA 是基於以下觀察:許多 Transformer 算子可以透過代數重新參數化。CODA 不將它們視為獨立的步驟,而是將它們整合到 GEMM 過程中。

運作方式

CODA 固定了 GEMM 主迴圈(mainloop)——矩陣乘法中計算密集的部分——並暴露了一組可組合的 epilogue 原語(primitives)。這些原語允許進行:

  • Scaling: 調整輸出分塊(tiles)的大小。
  • Reductions: 執行加總或其他歸約操作。
  • Pairwise Transformations: 應用逐元素函數(例如激活函數)。
  • Accumulation: 將結果累加到現有的張量中(例如殘差連接)。

透過在 GEMM 輸出分塊仍保留在晶片上的同時執行這些操作,CODA 避免了寫入和讀取全域記憶體的昂貴過程。其結果是一個既能保持專家級寫作的 GEMMs 性能,又能提供高階抽象穩健性的 kernel。

超越手動優化:LLM 的角色

CODA 框架最顯著的面向之一是其易用性。研究人員發現,無論是人類編寫的還是 LLM 編寫的 CODA kernel,都能達到高效率。這表明了低階 GPU kernel 編寫方式的轉轉變。

雖然 LLM 通常在處理低階硬體優化的複雜細節方面(例如管理共享記憶體 bank 或優化 warp-level 原語)感到吃力,但它們在處理高階組合方面表現出色。透過提供一個受限且可組合的 API,CODA 允許 LLM 將專家級寫作的區塊「黏合」在一起,形成一個功能完備且高效的 kernel。

正如社群成員在論文討論中所提到的:

"設計具有受限且可組合 API 的編譯器抽象,以便讓 LLM 可以輕鬆地將專家級寫作的區塊黏合在一起,這是一個明智之舉。我懷疑這最終會成為我們在邁向代理式開發(agentic development)時,代碼生成(codegens)的常態。"

結論

CODA 代表了一條將框架級生產力與硬體級效率結合起來的實用路徑。透過將 Transformer 區塊視為一系列 GEMM-epilogue 程式,它減少了記憶體瓶頸,並降低了記憶體頻寬的限制,並提供了一個結構化的環境,讓自動化工具——以及人類——可以快速迭代 kernel 性能。隨著訓練更大規模的模型持續推動記憶體頻寬的極限,像 CODA 這樣的抽象層將對於優化下一代 AI 基礎設施至關重要。

Sources