FlashQLA:適用於 CP / 反向傳播的融合線性注意力核心(針對 GDN)

Qwen 已發布 FlashQLA,一個高效能的線性注意力核心庫,旨在優化門控增量網路(Gated Delta Network,簡稱 GDN)層。FlashQLA 基於 TileLang 架構,在 NVIDIA Hopper GPU 上相較於 FLA Triton 核心,實現了 2‑3 倍的前向加速與 2 倍的反向加速,特別有利於預訓練與邊緣端代理推理。

最佳化 GDN 分段預填

FlashQLA 解決了原始 FLA 實作中 Gated Delta Network(GDN)分段預填(Chunked Prefill)流程的兩大效率瓶頸:

  1. 記憶體受限核心:標準流程會反覆讀寫中間變數 ($W, U, S$) 至高頻寬記憶體(HBM),產生大量開銷。
  2. GPU 使用率低:狀態空間模型(State Space Model,SSM)狀態的遞迴特性限制了同時執行的執行緒區塊數量為 batch_size * num_heads。在模型規模小、批次尺寸小或使用張量平行(Tensor Parallelism,TP)的情況下,會導致 GPU 流式多處理器(SM)閒置。

為了解決這些相互衝突的問題,FlashQLA 迴避了完整融合的核心(在小批次情境下會失效),改為將前向計算分割為兩個融合核心,並在其間插入上下文平行(Context Parallelism,CP)前處理步驟。

主要技術創新

門控驅動的自動卡內上下文平行(AutoCP)

FlashQLA 在張量平行(TP)、長序列與少頭數設定下實作自動卡內 CP 機制,以提升 SM 的利用率。它使用數學模型來決定最佳平行度 ($L = \lambda \sqrt{N}$),其中 $N$ 為分段數量,$L$ 為每個 CP 排的分段數。

為進一步降低開銷,FlashQLA 利用 GDN 門的指數衰減特性。對於門值 $\alpha_i \in (0,1)$ 的頭部,先前狀態的影響會指數衰減。FlashQLA 採用「熱身」過程(通常為 6–8 個分段)將狀態誤差降至噪聲底以下,從而省去昂貴的校正項 $M$ 矩陣計算,直接取得精確的子序列 $S_0$。

TileLang Warp 專用核心

FlashQLA 使用 TileLang 實作 warpgroup 專用核心。此架構在同一個 SM 內配置一個生產者 warpgroup 與三個消費者 warpgroup,透過共享記憶體交換資料,並以 mbarriers 同步。

  • 前向傳遞:三個消費者 warpgroup 分別計算 $V'、S、O$,採用 ping‑pong 結構以重疊計算與記憶體流量。
  • CP 前處理:單一融合核心同時處理原始的 $M$ 與 $S$ 計算,以及較輕量的滑動視窗熱身方法。
  • 反向傳遞:FlashQLA 將 bwd_dvbwd_dhubwd_dqkwgbwd_wy 融合為單一核心。由於晶片內資源受限,它依賴長計算鏈來隱藏記憶體流量,而非多階段流水線。

效能基準測試

在 NVIDIA H200 GPU 上相較於 FLA Triton 與 FlashInfer 基線進行的效能測試顯示出顯著提升,尤其隨著張量平行(TP)程度提升時更為明顯。

Model / TP Seqlen $h_{qk}$ $h_v$ FlashQLA FlashInfer FLA vs FLA vs FI
397B/122B TP8 1x32768 2 8 0.310ms 1.653ms 2.95×
397B/122B TP4 1x32768 4 16 0.486ms 1.654ms 2.57×
27B TP2 1x32768 8 24 0.659ms 1.616ms 2.37×
2B/0.8B TP1 1x32768 16 16 0.493ms 1.640ms 2.60×

實作與需求

FlashQLA 提供與 FLA 簽名相容的高階 API,以及前向與反向傳遞的低階入口點。

系統需求:

  • 硬體:NVIDIA SM90(Hopper)
  • 軟體:CUDA 12.8+、PyTorch 2.8+

程式碼與效能測試可於 github.com/QwenLM/FlashQLA 取得。


SUMMARY Qwen 已開源 FlashQLA,這是一個基於 TileLang 的高效能線性注意力核心庫,能在 NVIDIA Hopper GPU 上為門控增量網路(GDN)層提供 2‑3 倍的前向與 2 倍的反向加速。

TITLE FlashQLA:適用於 CP / 反向傳播的融合線性注意力核心(針對 GDN)

Sources