FlashAttentionFlashAttention 精確分塊注意力實作
白話解釋
透過分塊處理與逐步更新 softmax 統計量,避免將完整注意力中間矩陣存入大型記憶體的實作。反向傳播時會重算所需區塊,以換取較少的儲存與讀寫。
本文用法
本文將其作為直接基線;第二版繼承其精確計算與低中間記憶體需求,再改進運算及 GPU 工作分派。
原始 PDF 第 1 頁
GEMM矩陣乘法運算
白話解釋
執行矩陣相乘的高效能基本運算,GPU 可使用專門的運算單元加速。它在本文是衡量硬體計算能力的參照,而不是一套完整的注意力方法。
本文用法
作者以最佳化 GEMM 的峰值利用率對照原版 FlashAttention,指出注意力實作仍有改善空間;論文未明寫此縮寫的完整英文展開。
原始 PDF 第 1 頁
occupancyGPU 資源使用率(本文語境)
白話解釋
本文用它描述排入的平行工作是否足以使用 GPU 的運算資源。若可同時執行的 thread block 太少,即使單一區塊仍有大量計算,部分資源也可能閒置。
本文用法
作者將長序列、小 batch 或 head 數較少時的低 occupancy,列為沿序列維度增加 thread block 的理由;這裡採用作者所述的資源使用比例意義。
原始 PDF 第 2 頁
FlashAttention-2FlashAttention 第二版
白話解釋
本文提出的精確注意力 GPU 實作,沿用分塊計算,但調整非矩陣乘法工作、跨 thread block 平行化及 block 內 warp 分工。其目標是提升吞吐量,而非近似注意力結果。
本文用法
是全文的方法主體;評估時與原版 FlashAttention、Triton 與 xformers 等實作比較。
原始 PDF 第 1 頁
model FLOPs utilization模型 FLOPs 利用率
白話解釋
將模型訓練的估計浮點運算量與實際耗時換算為吞吐量,再相對於硬體理論峰值表示的比例。其值會受到所採 FLOPs 計數公式影響。
本文用法
本文以 72% 描述表 1 最高的每張 A100 訓練吞吐量;其整體訓練公式對 causal attention 的運算量未折半。
原始 PDF 第 1 頁
Tensor CoresTensor Cores 專用矩陣運算單元
白話解釋
GPU 上用來加速特定低精度矩陣乘法的專用硬體。它使矩陣乘法的理論吞吐量可能遠高於一般浮點運算。
本文用法
作者以其說明為何值得減少非矩陣乘法 FLOPs;H100 的第四代 Tensor Cores 則屬本文尚未特別運用的新功能。
原始 PDF 第 2 頁
HBM高頻寬記憶體
全名:high bandwidth memory
白話解釋
GPU 用於儲存輸入、輸出及其他較大型資料的記憶體。與晶片內的 shared memory 相比,頻繁傳輸大型中間矩陣可能造成顯著時間成本。
本文用法
標準注意力將 S、P 寫入 HBM;FlashAttention 系列盡量讓這些中間值留在晶片內處理。
原始 PDF 第 2 頁
SRAM晶片內 SRAM/shared memory
白話解釋
本文所說的晶片內儲存空間,供 thread block 暫放資料並讓 warp 交換中間值。容量較小,因此演算法需將輸入切成能在晶片內處理的區塊。
本文用法
原版 FlashAttention 用它減少 HBM 讀寫;第二版也針對 warp 間經由 shared memory 的多餘讀寫加以最佳化。論文未明寫 SRAM 縮寫的英文全名。
原始 PDF 第 2 頁
thread blocks執行緒區塊
白話解釋
GPU 將多個執行緒組成的排程單位,內部可再分成多個 warp。不同區塊可分別處理注意力矩陣的不同區域。
本文用法
FlashAttention-2 讓前向的不同列區塊、反向的不同欄區塊由不同 thread block 處理,以增加序列維度上的平行工作。
原始 PDF 第 2 頁
SMs串流多處理器
全名:streaming multiprocessors
白話解釋
GPU 上執行 thread block 的處理單元。若可排程的 block 數量不足,便難以同時使用多個處理單元。
本文用法
作者以 A100 的處理單元數說明,只有 batch 與 head 維度的平行工作時,長序列、小 batch 可能無法充分使用 GPU。
原始 PDF 第 2 頁
warpswarp 執行緒群組
白話解釋
本文所述由 32 個執行緒組成的群組;同一 warp 內的執行緒可快速協作。不同 warp 若要交換某些中間結果,可能需要透過 shared memory。
本文用法
第 3.3 節在單一 thread block 內重新分配各 warp 負責的 Q、K、V 區域,是第二版減少通訊的重要設計。
原始 PDF 第 2 頁
tiling分塊運算
白話解釋
將大型矩陣運算拆成較小區塊,逐塊載入與計算。若區塊可暫留於晶片內記憶體,就能避免頻繁搬動完整中間矩陣。
本文用法
是 FlashAttention 系列的既有基礎;FlashAttention-2 進一步調整這些區塊如何分派給 thread block 與 warp。
原始 PDF 第 3 頁
online softmax線上 softmax
白話解釋
逐區塊更新逐列最大值與指數和,並依新統計量重新縮放先前累積結果的方法。輔助理解:即使一列的各段不能同時放在記憶體中,也能逐段處理並得到整列的正規化結果。
本文用法
使分塊注意力無須儲存完整分數與權重矩陣仍可計算相同輸出;第二版再調整正規化的時機。
原始 PDF 第 3 頁
logsumexp對數指數和
白話解釋
本文使用的逐列 softmax 統計量 \(L=m+\log(\ell)\),其中 \(m\) 是逐列最大值,\(\ell\) 是穩定化後的指數和。它將反向傳播所需的兩項統計資訊合併保存。
本文用法
Algorithm 1 輸出 \(L\),Algorithm 2 利用它重建 softmax 權重;這裡不是額外儲存整個注意力矩陣。
原始 PDF 第 5 頁
causal mask因果遮罩
白話解釋
自迴歸注意力中禁止位置 \(i\) 讀取未來位置 \(j>i\) 的規則。在分塊計算時,若整個區塊都位於被禁止的區域,便能跳過該區塊。
本文用法
本文以有、無 causal mask 分別測試吞吐量;注意力基準測試的有遮罩 FLOPs 計數約折半,整體訓練計數則沿用未折半的文獻公式。
原始 PDF 第 6 頁
MQA多查詢注意力
全名:Multi-query attention
白話解釋
讓多個 query head 共用同一組 key、value head 的注意力變體。這可降低推論時須保存的 key/value 資料量。
本文用法
作者說明實作可透過 head 索引處理這種共用關係,反向傳播則須加總相應的 key、value 梯度。
原始 PDF 第 7 頁
GQA分組查詢注意力
全名:grouped-query attention
白話解釋
將 query head 分組,使同組的多個 query head 共用 key、value head 的注意力變體。它與 MQA 同屬減少推論期間 key/value 快取大小的設計。
本文用法
本文將其與 MQA 一併列為支援的變體,並指出反向傳播須彙整共用 head 所對應的梯度。
原始 PDF 第 7 頁
atomic adds原子加法
白話解釋
多個平行工作單位更新同一資料位置時,可安全累加貢獻的操作。它避免某個更新覆蓋另一個更新,但仍代表需要跨工作單位協調。
本文用法
反向傳播依欄區塊分派 thread block 時,各 block 對同一 dQ 的貢獻以 atomic adds 累加。
原始 PDF 第 8 頁
split-K依 K 切分的 warp 工作配置
白話解釋
本文用來稱呼原版前向傳播中,將 K、V 分給不同 warp、再合併部分輸出的配置。由於多個 warp 共同形成結果,須交換並加總中間值。
本文用法
FlashAttention-2 改切分 Q,讓各 warp 產生各自的輸出片段,以避免原配置的前向跨 warp 歸約。
原始 PDF 第 9 頁
register spilling暫存器溢出
白話解釋
運算所需的暫存器超過可用資源時,部分資料無法繼續保留在暫存器中的情況。這可能引入較慢的存取,使原本想藉增大區塊取得的效益消失。
本文用法
作者將它列為 block 大小不能一味增加的原因之一,因此依 head dimension 與 GPU 記憶體資源手動調整區塊。
原始 PDF 第 9 頁