RESEARCH NOTE 雙語閱讀 · 重點批註

FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning

Tri DaoarXiv preprint, arXiv:2307.08691v12023

arXiv:2307.08691v1

模型:gpt-6-sol · REASONING:xhigh · DETAIL:high
輸入:32,573 tokens · 輸出:27,197 tokens(思考:9,120 tokens、其他輸出:18,077 tokens) · 總計:59,770 tokens

一句話掌握THE TAKEAWAY

FlashAttention-2 在維持精確注意力計算與線性額外記憶體需求的前提下,重新安排演算法、thread block 與 warp 的工作,提升 GPU 吞吐量。

研究問題

為何已減少記憶體讀寫的 FlashAttention 仍遠慢於高效矩陣乘法?如何在長序列、小 batch 的情境下進一步加速前向與反向計算?

動機

原版 FlashAttention 雖比標準注意力快,但作者觀察到 GPU 資源使用不足,以及 warp 間不必要的 shared memory 通訊;序列越長、batch 越小,平行度問題越突出。

方法

減少非矩陣乘法運算;以前向的 query 列區塊、反向的 key/value 欄區塊增加 thread block 平行度;前向在 block 內切分 Q 而非 K、V,以避免 warp 間歸約。

主要貢獻

  • 保留分塊 online softmax 的精確輸出,調整正規化與反向傳播所需的統計量。
  • 沿序列維度分派前向及反向工作,改善長序列且 batch 或 head 數較少時的 GPU 資源使用率。
  • 改變 warp 間工作分配,減少經由 shared memory 交換中間結果。

主要結果

  • A100 的注意力基準測試中,FlashAttention-2 約比 FlashAttention 快 2 倍;依測試設定不同,前向吞吐量最高達理論峰值的 73%,反向最高達 63%。
  • 在 8 張 A100 上訓練所測試的 GPT-style 模型時,最高達每張 GPU 225 TFLOPs/s;相較無 FlashAttention 的基線,報告的最大加速為 2.8 倍。

限制

  • block 大小仍須依 head dimension 與裝置記憶體手動調整;過大會造成 register spilling,甚至無法執行。
  • 論文中的 H100 測試尚未使用 TMA 或第四代 Tensor Cores 等新硬體功能;進一步加速屬作者預期,並非已驗證結果。
  • 訓練吞吐量沿用文獻中的 FLOPs 計算式,對 causal attention 不折半;閱讀其 TFLOPs/s 與利用率時須留意計數方式。
重要度
回到頂端 ↑

1. 研究問題與效能缺口

長上下文使注意力的時間與記憶體成本隨序列長度平方成長。本文的直接研究缺口不是近似注意力,而是原版 FlashAttention 雖已減少記憶體讀寫,其 GPU 工作分配仍未充分發揮硬體吞吐量。

p. 1

原文 SOURCE

Scaling up the context length of Transformers [18] is a challenge, since the attention layer at their heart has runtime and memory requirements quadratic in the input sequence length. Ideally, we would like to go beyond the standard 2k sequence length limit to train models to understand books, high resolution images, and long-form videos. Just within the last year, there have been several language models with much longer context than before: GPT-4 [12] with context length 32k, MosaicML’s MPT with context length 65k, and Anthropic’s Claude with context length 100k. Emerging use cases such as long document querying and story writing have demonstrated a need for models with such long context.

繁體中文 TRANSLATION

擴大 Transformer [18] 的上下文長度是一項挑戰,因為其核心注意力層的執行時間與記憶體需求都隨輸入序列長度呈平方成長。理想情況下,我們希望超越標準的 2k 序列長度限制,訓練能理解書籍、高解析度影像及長篇影片的模型。就在過去一年,已有數個語言模型提供比以往長得多的上下文:GPT-4 [12] 的上下文長度為 32k、MosaicML 的 MPT 為 65k,Anthropic 的 Claude 為 100k。長文件查詢與故事寫作等新興用途,顯示模型確實需要如此長的上下文。

p. 2

原文 SOURCE

However, context length increases even more, FlashAttention is still not nearly as efficient as other primitives such as matrix-multiply (GEMM). In particular, while FlashAttention is already 2-4× faster than a standard attention implementation, the forward pass only reaches 30-50% of the theoretical maximum FLOPs/s of the device (Fig. 5), while the backward pass is even more challenging, reaching only 25-35% of maximum throughput on A100 GPU (Fig. 6). In contrast, optimized GEMM can reach up to 80-90% of the theoretical maximum device throughput. Through careful profiling, we observe that FlashAttention still has suboptimal work partitioning between different thread blocks and warps on the GPU, causing either low-occupancy or unnecessary shared memory reads/writes.

繁體中文 TRANSLATION

然而,隨著上下文長度進一步增加,FlashAttention 的效率仍遠不及矩陣乘法(GEMM)等其他基本運算。具體而言,FlashAttention 雖然已比標準注意力實作快 2–4 倍,其前向傳播仍僅達裝置理論最高 FLOPs/s 的 30–50%(圖 5);反向傳播更具挑戰性,在 A100 GPU 上僅達最高吞吐量的 25–35%(圖 6)。相較之下,經過最佳化的 GEMM 可達裝置理論最高吞吐量的 80–90%。透過仔細分析效能,我們發現 FlashAttention 在 GPU 的不同 thread block 與 warp 之間仍有不理想的工作分配,導致資源使用率偏低,或產生不必要的 shared memory 讀寫。

2. 必要背景:記憶體讀寫與精確分塊注意力

標準注意力須將大型中間矩陣寫入 HBM;原版 FlashAttention 以分塊、online softmax 與反向重算避免這些讀寫。理解這個既有基礎,才能辨認 FlashAttention-2 改進的是哪一層成本。

p. 3

原文 SOURCE

Standard attention implementations materialize the matrices S and P to HBM, which takes \(O(N^2)\) memory. Often \(N \gg d\) (typically \(N\) is on the order of 1k–8k and \(d\) is around 64–128). The standard attention implementation (1) calls the matrix multiply (GEMM) subroutine to multiply \(S = QK^\top\), writes the result to HBM, then (2) loads \(S\) from HBM to compute softmax and write the result \(P\) to HBM, and finally (3) calls GEMM to get \(O = PV\). As most of the operations are bounded by memory bandwidth, the large number of memory accesses translates to slow wall-clock time. Moreover, the required memory is \(O(N^2)\) due to having to materialize S and P. Moreover, one has to save \(P \in \mathbb{R}^{N \times N}\) for the backward pass to compute the gradients.

繁體中文 TRANSLATION

標準注意力實作會將矩陣 S 與 P 實際寫入 HBM,需要 \(O(N^2)\) 記憶體。通常 \(N \gg d\)(\(N\) 一般約為 1k–8k,而 \(d\) 約為 64–128)。標準實作會先(1)呼叫矩陣乘法(GEMM)子程序計算 \(S = QK^\top\),並將結果寫入 HBM;接著(2)從 HBM 載入 \(S\),計算 softmax,再將結果 \(P\) 寫入 HBM;最後(3)呼叫 GEMM 得到 \(O = PV\)。由於大部分操作受記憶體頻寬限制,大量記憶體存取使實際執行時間變長。此外,實際儲存 S 與 P 使記憶體需求達 \(O(N^2)\)。反向傳播計算梯度時,還必須保留 \(P \in \mathbb{R}^{N \times N}\)。

p. 3

原文 SOURCE

FlashAttention applies the classical technique of tiling to reduce memory IOs, by (1) loading blocks of inputs from HBM to SRAM, (2) computing attention with respect to that block, and then (3) updating the output without writing the large intermediate matrices S and P to HBM. As the softmax couples entire rows or blocks of row, online softmax [11, 13] can split the attention computation into blocks, and rescale the output of each block to finally get the right result (with no approximation). By significantly reducing the amount of memory reads/writes, FlashAttention yields 2-4× wall-clock speedup over optimized baseline attention implementations.

繁體中文 TRANSLATION

FlashAttention 運用經典的分塊(tiling)技術減少記憶體輸入輸出:(1)將輸入區塊從 HBM 載入 SRAM,(2)針對該區塊計算注意力,再(3)更新輸出,而不把大型中間矩陣 S 與 P 寫入 HBM。由於 softmax 會將整列或列區塊中的元素耦合在一起,online softmax [11, 13] 可以把注意力計算拆成多個區塊,並重新縮放各區塊的輸出,最後得到正確結果,而不作近似。記憶體讀寫量大幅降低後,FlashAttention 相較經過最佳化的基線注意力實作,實際執行時間縮短至約原本的四分之一至二分之一。

Figure 1原始 PDF 第 4 頁

Diagram of how FlashAttention forward pass is performed, when the key K is partitioned into two blocks and the value V is also partitioned into two blocks. By computing attention with respect to each block and rescaling the output, we get the right answer at the end, while avoiding expensive memory reads/writes of the intermediate matrices S and P. We simplify the diagram, omitting the step in softmax that subtracts each element by the row-wise max.FlashAttention 前向傳播示意圖,其中 key K 分為兩個區塊,value V 也分為兩個區塊。透過逐區塊計算注意力並重新縮放輸出,最後可得到正確答案,同時避免對中間矩陣 S 與 P 進行昂貴的記憶體讀寫。圖中為求簡化,省略 softmax 將每個元素減去逐列最大值的步驟。

這張圖在說什麼

圖中示範同一批 query 依序與兩個 K、V 區塊計算,再將先前的輸出重新縮放後合併。大型分數與權重中間值留在 SRAM 計算,不必完整寫回 HBM。

怎麼看

先沿 Q、兩個 K 區塊到分數區塊,再沿指數運算、兩個 V 區塊讀到右側輸出;最後看右側標示的重新縮放。藍色虛線框代表儲存在 HBM 的資料,橘色虛線框代表在 SRAM 計算而未寫入 HBM 的中間值;圖中刻意省略逐列減最大值,不能將示意式當作完整的數值穩定 softmax。

p. 5

原文 SOURCE

In the backward pass, by re-computing the values of the attention matrices S and P once blocks of inputs Q, K, V are already loaded to SRAM, FlashAttention avoids having to store large intermediate values. By not having to save the large matrices S and P of size \(N \times N\), FlashAttention yields 10-20× memory saving depending on sequence length (memory required in linear in sequence length \(N\) instead of quadratic). The backward pass also achieves 2-4× wall-clock speedup due to reduce memory reads/writes.

繁體中文 TRANSLATION

在反向傳播時,FlashAttention 於輸入 Q、K、V 的區塊已載入 SRAM 後,重新計算注意力矩陣 S 與 P 的值,因此不必儲存大型中間值。由於不需要保存大小為 \(N \times N\) 的 S 與 P,FlashAttention 依序列長度而定,可節省 10–20 倍記憶體(記憶體需求改為隨序列長度 \(N\) 線性成長,而非平方成長)。反向傳播也因記憶體讀寫減少,取得 2–4 倍的實際執行速度提升。

3.1 演算法調整:計算配比、統計量與遮罩

GPU 的矩陣乘法吞吐量遠高於一般運算,因此作者減少 softmax 相關的非矩陣乘法工作;前向逐區塊累積、最後才正規化,反向只保存逐列 logsumexp。方法仍計算不近似的注意力輸出。

p. 5

原文 SOURCE

We tweak the algorithm from FlashAttention to reduce the number of non-matmul FLOPs. This is because modern GPUs have specialized compute units (e.g., Tensor Cores on Nvidia GPUs) that makes matmul much faster. As an example, the A100 GPU has a max theoretical throughput of 312 TFLOPs/s of FP16/BF16 matmul, but only 19.5 TFLOPs/s of non-matmul FP32. Another way to think about this is that each non-matmul FLOP is 16× more expensive than a matmul FLOP. To maintain high throughput (e.g., more than 50% of the maximum theoretical TFLOPs/s), we want to spend as much time on matmul FLOPs as possible.

繁體中文 TRANSLATION

我們調整 FlashAttention 演算法,以減少非矩陣乘法的浮點運算次數。原因是現代 GPU 具備專用運算單元(例如 Nvidia GPU 的 Tensor Cores),使矩陣乘法快得多。以 A100 GPU 為例,FP16/BF16 矩陣乘法的理論最高吞吐量為 312 TFLOPs/s,但非矩陣乘法的 FP32 運算僅為 19.5 TFLOPs/s。換個角度看,每一次非矩陣乘法浮點運算的成本相當於一次矩陣乘法運算的 16 倍。為維持高吞吐量(例如超過理論最高 TFLOPs/s 的 50%),我們希望盡可能把時間用在矩陣乘法運算上。

p. 5

原文 SOURCE

We do not have to save both the max \(m^{(j)}\) and the sum of exponentials \(\ell^{(j)}\) for the backward pass. We only need to store the logsumexp \(L^{(j)} = m^{(j)} + \log(\ell^{(j)})\).

繁體中文 TRANSLATION

反向傳播不必同時儲存最大值 \(m^{(j)}\) 與指數和 \(\ell^{(j)}\)。只需要儲存 logsumexp:\(L^{(j)} = m^{(j)} + \log(\ell^{(j)})\)。

p. 6

原文 SOURCE

One common use case of attention is in auto-regressive language modeling, where we need to apply a causal mask to the attention matrix S (i.e., any entry \(S_{ij}\) with \(j > i\) is set to \(-\infty\)). As FlashAttention and FlashAttention-2 already operate by blocks, for any blocks where all the column indices are more than the row indices (approximately half of the blocks for large sequence length), we can skip the computation of that block. This leads to around 1.7-1.8× speedup compared to attention without the causal mask.

繁體中文 TRANSLATION

注意力的一項常見用途是自迴歸語言模型,此時必須對注意力矩陣 S 套用 causal mask(亦即凡 \(j > i\) 的元素 \(S_{ij}\) 都設為 \(-\infty\))。FlashAttention 與 FlashAttention-2 本來就以區塊為單位運作;若某區塊的所有欄索引均大於列索引,對長序列而言約占一半區塊,便可跳過該區塊的計算。因此相較於不使用 causal mask 的注意力,速度可提升約 1.7–1.8 倍。

p. 7

原文 SOURCE

As with FlashAttention, Algorithm 1 returns the correct output \(O = \operatorname{softmax}(QK^\top)V\) (with no approximation), using \(O(N^2d)\) FLOPs and requires \(O(N)\) additional memory beyond inputs and output (to store the logsumexp \(L\)). The proof is almost the same as the proof of Dao et al. [5, Theorem 1], so we omit it here.

繁體中文 TRANSLATION

如同 FlashAttention,演算法 1 會回傳正確輸出 \(O = \operatorname{softmax}(QK^\top)V\),不採用近似;其運算量為 \(O(N^2d)\) 浮點運算,除輸入與輸出之外,額外只需 \(O(N)\) 記憶體來儲存 logsumexp \(L\)。證明與 Dao 等人 [5, Theorem 1] 的證明幾乎相同,因此本文不再列出。

3.2 序列維度的 thread block 平行化

當 batch 與 head 數不足以填滿 GPU 時,作者將前向的 query 列區塊、反向的 key/value 欄區塊分派給不同 thread block;前向可彼此獨立,反向更新 dQ 則需要跨 block 累加。

p. 8

原文 SOURCE

We see that the outer loop (over sequence length) is embarrassingly parallel, and we schedule them on different thread blocks that do not need to communicate with each other. We also parallelize over the batch dimension and number of heads dimension, as done in FlashAttention. The increased parallelism over sequence length helps improve occupancy (fraction of GPU resources being used) when the batch size and number of heads are small, leading to speedup in this case.

繁體中文 TRANSLATION

我們發現外層迴圈(遍歷序列長度)很容易平行化,因此把各次迴圈分派給不同 thread block,彼此不必通訊。我們也和 FlashAttention 一樣,沿 batch 維度及 head 數量維度平行化。當 batch 大小與 head 數量較少時,序列維度增加的平行工作可提高 occupancy(GPU 資源的使用比例),進而在這種情況下加速。

p. 8

原文 SOURCE

Notice that the only shared computation between different column blocks is in update dQ in Algorithm 2, where we need to load \(\mathrm{d}Q_i\) from HBM to SRAM, then on chip, update \(\mathrm{d}Q_i \leftarrow \mathrm{d}Q_i + \mathrm{d}S_i^{(j)}K_j\), and write back to HBM. We thus parallelize over the sequence length dimension as well, and schedule 1 thread block for each column block of the backward pass. We use atomic adds to communicate between different thread blocks to update dQ.

繁體中文 TRANSLATION

請注意,演算法 2 中不同欄區塊唯一會共享的計算是更新 dQ:我們必須把 \(\mathrm{d}Q_i\) 從 HBM 載入 SRAM,在晶片內執行 \(\mathrm{d}Q_i \leftarrow \mathrm{d}Q_i + \mathrm{d}S_i^{(j)}K_j\),再寫回 HBM。因此,我們也沿序列長度維度平行化反向傳播,讓每個欄區塊各由一個 thread block 處理。我們使用 atomic adds 讓不同 thread block 共同更新 dQ。

Figure 2原始 PDF 第 8 頁

In the forward pass (left), we parallelize the workers (thread blocks) where each worker takes care of a block of rows of the attention matrix. In the backward pass (right), each worker takes care of a block of columns of the attention matrix.前向傳播(左)將工作者(thread block)平行化,每個工作者負責注意力矩陣的一個列區塊。反向傳播(右)則由每個工作者負責注意力矩陣的一個欄區塊。

這張圖在說什麼

左圖的水平色帶表示前向工作依列區塊分派;右圖的垂直色帶表示反向工作依欄區塊分派。圖示強調的是 thread block 之間如何切分工作,而非單一 block 內的 warp 分工。

怎麼看

先看左圖:同色的列區塊由同一 worker 負責;再看右圖:同色的欄區塊由同一 worker 負責。白色區域對應圖示中的未計算區塊;此圖沒有速度座標軸,反向更新 dQ 的 atomic adds 須配合旁邊文字理解。

3.3 warp 分工與 block 大小的取捨

FlashAttention-2 在單一 thread block 內改切分 Q,使各 warp 產生各自的輸出列片段,減少原本切分 K、V 時所需的 shared memory 歸約。較大的計算區塊雖能減少讀寫,卻受暫存器與 shared memory 容量限制。

p. 9

原文 SOURCE

For each block, FlashAttention splits K and V across 4 warps while keeping Q accessible by all warps. Each warp multiplies to get a slice of \(QK^\top\), then they need to multiply with a slice of V and communicate to add up the result. This is referred to as the “split-K” scheme. However, this is inefficient since all warps need to write their intermediate results out to shared memory, synchronize, then add up the intermediate results. These shared memory reads/writes slow down the forward pass in FlashAttention. In FlashAttention-2, we instead split Q across 4 warps while keeping K and V accessible by all warps. After each warp performs matrix multiply to get a slice of \(QK^\top\), they just need to multiply with their shared slice of V to get their corresponding slice of the output. There is no need for communication between warps. The reduction in shared memory reads/writes yields speedup (Section 4).

繁體中文 TRANSLATION

對每個區塊,FlashAttention 將 K 與 V 分散到 4 個 warp,同時讓所有 warp 都能存取 Q。每個 warp 進行矩陣乘法,得到 \(QK^\top\) 的一部分,再與 V 的一部分相乘,並透過通訊將結果相加。這稱為「split-K」配置。然而,這種方式效率不佳,因為所有 warp 都得將中間結果寫入 shared memory、進行同步,然後加總中間結果。這些 shared memory 讀寫拖慢了 FlashAttention 的前向傳播。 在 FlashAttention-2 中,我們改將 Q 分散到 4 個 warp,並讓所有 warp 都能存取 K 與 V。每個 warp 經矩陣乘法取得 \(QK^\top\) 的一部分後,只需再與各 warp 均可存取的 V 相乘,即可得到自己負責的輸出片段。warp 之間不需要通訊。shared memory 讀寫的減少帶來了加速(第 4 節)。

Figure 3原始 PDF 第 9 頁

Work partitioning between different warps in the forward pass前向傳播中不同 warp 之間的工作分配

這張圖在說什麼

左圖的 FlashAttention 把 K、V 分散到不同 warp,需合併各 warp 的部分結果;右圖的 FlashAttention-2 改分散 Q,使各 warp 各自產生對應的輸出片段。這支撐了作者關於減少跨 warp shared memory 讀寫的主張。

怎麼看

先看圖例:藍色代表所有 warp 均可存取,橘色代表分配到不同 warp。接著比較 (a) 中藍色的 Q、橘色的 K 與 V,及 (b) 中橘色的 Q、藍色的 K 與 V;沿 \(QK^\top\) 再乘 V 的資料流,判斷是否須將多個 warp 的結果相加。

p. 9

原文 SOURCE

Increasing block sizes generally reduces shared memory loads/stores, but increases the number of registers required and the total amount of shared memory. Past a certain block size, register spilling causes significant slowdown, or the amount of shared memory required is larger than what the GPU has available, and the kernel cannot run at all. Typically we choose blocks of size \(\{64, 128\} \times \{64, 128\}\), depending on the head dimension \(d\) and the device shared memory size. We manually tune for each head dimensions since there are essentially only 4 choices for block sizes, but this could benefit from auto-tuning to avoid this manual labor. We leave this to future work.

繁體中文 TRANSLATION

增大區塊通常可減少 shared memory 的載入與儲存,但會增加所需暫存器數量及 shared memory 總量。區塊大到某個程度後,register spilling 會導致顯著變慢;或者所需 shared memory 超過 GPU 可提供的容量,使 kernel 完全無法執行。我們通常依 head dimension \(d\) 與裝置的 shared memory 容量,選用 \(\{64, 128\} \times \{64, 128\}\) 的區塊大小。 由於區塊大小基本上只有 4 種選擇,我們針對各個 head dimension 手動調整;若使用自動調校,便可免去這項人工工作。我們將其留待未來研究。

4. 實驗結果與吞吐量計數方式

作者分別量測 A100 上的前向、反向及合併注意力吞吐量,並在 8 張 A100 上比較 GPT-style 模型的訓練速度。兩種實驗的 FLOPs 計數方式不同,尤其要注意 causal mask 是否使計數折半。

p. 10

原文 SOURCE

Benchmark setting: we vary the sequence length from 512, 1k, ..., 16k, and set batch size so that the total number of tokens is 16k. We set hidden dimension to 2048, and head dimension to be either 64 or 128 (i.e., 32 heads or 16 heads). To calculate the FLOPs of the forward pass, we use: \[4 \cdot \text{seqlen}^{2} \cdot \text{head dimension} \cdot \text{number of heads}.\] With causal mask, we divide this number by 2 to account for the fact that approximately only half of the entries are calculated. To get the FLOPs of the backward pass, we multiply the forward pass FLOPs by 2.5 (since there are 2 matmuls in the forward pass and 5 matmuls in the backward pass, due to recomputation).

繁體中文 TRANSLATION

基準測試設定如下:我們將序列長度從 512、1k 一直變動到 16k,並調整 batch 大小,使 token 總數保持在 16k。我們將 hidden dimension 設為 2048,head dimension 設為 64 或 128(亦即 32 個或 16 個 head)。計算前向傳播 FLOPs 時,採用: \[4 \cdot \text{seqlen}^{2} \cdot \text{head dimension} \cdot \text{number of heads}.\] 若使用 causal mask,我們將此數值除以 2,因為實際計算的元素大約只有一半。計算反向傳播 FLOPs 時,則將前向傳播的 FLOPs 乘以 2.5(因重算之故,前向有 2 次矩陣乘法,反向有 5 次)。

p. 10

原文 SOURCE

We measure the runtime of different attention methods on an A100 80GB SXM4 GPU for different settings (without / with causal mask, head dimension 64 or 128). We report the results in Fig. 4, Fig. 5 and Fig. 6, showing that FlashAttention-2 is around 2× faster than FlashAttention and FlashAttention in xformers (the “cutlass” implementation). FlashAttention-2 is around 1.3-1.5× faster than FlashAttention in Triton in the forward pass and around 2× faster in the backward pass. Compared to a standard attention implementation in PyTorch, FlashAttention-2 can be up to 10× faster.

繁體中文 TRANSLATION

我們在 A100 80GB SXM4 GPU 上,針對不同設定(有無 causal mask、head dimension 為 64 或 128)量測各種注意力方法的執行時間。我們在圖 4、圖 5 與圖 6 報告結果:FlashAttention-2 約比 FlashAttention 及 xformers 中的 FlashAttention(「cutlass」實作)快 2 倍。與 Triton 中的 FlashAttention 相比,FlashAttention-2 的前向傳播約快 1.3–1.5 倍,反向傳播約快 2 倍。相較於 PyTorch 中的標準注意力實作,FlashAttention-2 最多可快 10 倍。

Figure 4原始 PDF 第 10 頁

Attention forward + backward speed on A100 GPUA100 GPU 上注意力前向與反向合併的速度

這張圖在說什麼

四個分圖比較有無 causal mask、head dimension 為 64 或 128 時的合併吞吐量。各設定下 FlashAttention-2 的柱狀值高於所列基線,支撐其整體注意力計算約有顯著加速的結論。

怎麼看

先依 (a)–(d) 確認遮罩與 head dimension,再讀橫軸序列長度、縱軸 TFLOPs/s;數值越高代表依文中 FLOPs 計法得到的吞吐量越高。對同一組序列長度比較紫色 FlashAttention-2、橘色 FlashAttention、紅色 Triton、綠色 xformers 與藍色 PyTorch;圖示 OOM 表示該設定的 PyTorch 無結果,不是速度為零。

Figure 5原始 PDF 第 11 頁

Attention forward speed on A100 GPUA100 GPU 上注意力前向傳播的速度

這張圖在說什麼

這張圖單獨檢驗前向傳播,對應作者關於減少非矩陣乘法工作及跨 warp 通訊的設計。FlashAttention-2 在測試設定下維持高於所列基線的前向吞吐量;圖中最高的紫色柱標示為 227 TFLOPs/s。

怎麼看

先分別閱讀無遮罩的 (a)(b) 與有遮罩的 (c)(d),再在固定 head dimension、序列長度下比較各色柱;縱軸 TFLOPs/s 越高越好。遮罩設定的 FLOPs 計數已約略折半,因此不要把不同分圖的柱高直接當成相同運算量下的延遲比。

Figure 6原始 PDF 第 12 頁

Attention backward speed on A100 GPUA100 GPU 上注意力反向傳播的速度

這張圖在說什麼

這張圖將反向傳播獨立呈現,可檢查反向欄區塊平行化是否也有收益。FlashAttention-2 的紫色柱在所示設定中高於其他方法;無遮罩、head dimension 128、序列長度 16k 的柱標示為 196 TFLOPs/s。

怎麼看

先按分圖確認遮罩與 head dimension,再沿序列長度比較紫色與其他顏色的柱;縱軸吞吐量越高越好。反向 FLOPs 採前向計數的 2.5 倍,故應優先在同一分圖、相同序列長度下比較方法。

p. 10

原文 SOURCE

When used end-to-end to train GPT-style models of size 1.3B and 2.7B on sequence lengths either 2k or 8k, FlashAttention-2 yields up to 1.3× speedup compared to FlashAttention and 2.8× speedup compared to a baseline without FlashAttention. FlashAttention-2 reaches up to 225 TFLOPs/s (72% model FLOPs utilization) per A100 GPU.

繁體中文 TRANSLATION

將 FlashAttention-2 用於端到端訓練參數量為 1.3B 或 2.7B、序列長度為 2k 或 8k 的 GPT-style 模型時,相較 FlashAttention 最多加速 1.3 倍,相較未使用 FlashAttention 的基線最多加速 2.8 倍。FlashAttention-2 在每張 A100 GPU 上最高達 225 TFLOPs/s(模型 FLOPs 利用率為 72%)。

Table 1原始 PDF 第 12 頁

Training speed (TFLOPs/s/GPU) of GPT-style models on 8×A100 GPUs. FlashAttention-2 reaches up to 225 TFLOPs/s (72% model FLOPs utilization). We compare against a baseline running without FlashAttention.在 8 張 A100 GPU 上訓練 GPT-style 模型的速度(TFLOPs/s/GPU)。FlashAttention-2 最高達 225 TFLOPs/s(模型 FLOPs 利用率為 72%)。比較對象包含未使用 FlashAttention 的基線。
ModelWithout FlashAttentionFlashAttentionFlashAttention-2
GPT3-1.3B 2k context142 TFLOPs/s189 TFLOPs/s196 TFLOPs/s
GPT3-1.3B 8k context72 TFLOPs/s170 TFLOPs/s220 TFLOPs/s
GPT3-2.7B 2k context149 TFLOPs/s189 TFLOPs/s205 TFLOPs/s
GPT3-2.7B 8k context80 TFLOPs/s175 TFLOPs/s225 TFLOPs/s

表格註記:原表以每張 GPU 的吞吐量呈現,實驗使用 8 張 A100;72% 模型 FLOPs 利用率出現在表說,而非每個儲存格。第 4.2 節正文的「compared to FlashAttention-2」為自我比較的筆誤;第 4 節較早敘述及本表顯示比較對象為原版 FlashAttention。

這張表在說什麼

四列涵蓋 1.3B、2.7B 模型各自的 2k、8k 上下文;三個速度欄分別代表無 FlashAttention、原版 FlashAttention 與第二版。FlashAttention-2 各列數值最高,其中 2.7B、8k 為每張 GPU 225 TFLOPs/s。

怎麼看

先選定同一模型及上下文長度的一列,再由左至右比較三個方法;單位均為每張 GPU 的 TFLOPs/s,越高越好。要觀察長上下文效果,可各自比較 8k 列中的基線與 FlashAttention-2;不可把不同模型列的差異全部歸因於注意力實作。

p. 11

原文 SOURCE

Note that we calculate the FLOPs by the formula, following Megatron-LM [16] (and many other papers and libraries): \[6 \cdot \text{seqlen} \cdot \text{number of params} + 12 \cdot \text{number of layers} \cdot \text{hidden dim} \cdot \text{seqlen}^{2}.\] The first term accounts for the FLOPs due to weight–input multiplication, and the second term accounts for the FLOPs due to attention. However, one can argue that the second term should be halved, as with causal mask we only need to compute approximately half the number of elements in attention. We choose to follow the formula from the literature (without dividing the attention FLOPs by 2) for consistency.

繁體中文 TRANSLATION

請注意,我們依循 Megatron-LM [16](及許多其他論文與函式庫)的公式計算 FLOPs: \[6 \cdot \text{seqlen} \cdot \text{number of params} + 12 \cdot \text{number of layers} \cdot \text{hidden dim} \cdot \text{seqlen}^{2}.\] 第一項是權重與輸入相乘所需的 FLOPs,第二項是注意力所需的 FLOPs。不過,也可以主張第二項應減半,因為使用 causal mask 時,注意力中大約只需計算一半元素。我們為了與既有文獻一致,選擇沿用該公式,不將注意力 FLOPs 除以 2。

5. 適用邊界與後續方向

H100 結果展示同一實作在另一款 GPU 上的表現,但作者明言尚未利用其新指令。跨裝置、新資料型別、自動化最佳化及與稀疏注意力結合,仍屬未來工作。

p. 11

原文 SOURCE

Just running the same implementation on H100 GPUs (using no special instructions to make use of new features such as TMA and 4th-gen Tensor Cores), we obtain up to 335 TFLOPs/s (Fig. 7). We expect that by using new instructions, we can obtain another 1.5x-2x speedup on H100 GPUs. We leave that to future work.

繁體中文 TRANSLATION

僅將同一實作執行於 H100 GPU,而未使用 TMA 或第四代 Tensor Cores 等新功能所需的特殊指令,我們便得到最高 335 TFLOPs/s(圖 7)。我們預期若使用新指令,在 H100 GPU 上還能再加速 1.5–2 倍。我們將此留待未來研究。

Figure 7原始 PDF 第 13 頁

Attention forward + backward speed on H100 GPUH100 GPU 上注意力前向與反向合併的速度

這張圖在說什麼

四個分圖在 H100 上比較 PyTorch、FlashAttention 與 FlashAttention-2,紫色柱在所示條件下最高。它展示同一實作在另一款 GPU 上的結果,並不是利用 H100 新指令後之預期加速的驗證。

怎麼看

先確認 (a)–(d) 的遮罩與 head dimension,再讀橫軸序列長度及縱軸 TFLOPs/s;越高越好。在相同設定下比較藍、橘、紫柱;(b) 的紫色柱在 8k 標 335、16k 標 338,引用數值時須區分圖上條件與正文的「up to 335」說法。OOM 表示 PyTorch 該條件未取得速度值。

p. 12

原文 SOURCE

In the near future, we plan to collaborate with researchers and engineers to make FlashAttention widely applicable in different kinds of devices (e.g., H100 GPUs, AMD GPUs), as well as new data types such as FP8. As an immediate next step, we plan to optimize FlashAttention-2 for H100 GPUs to use new hardware features (TMA, 4th-gen Tensor Cores, fp8). Combining the low-level optimizations in FlashAttention-2 with high-level algorithmic changes (e.g., local, dilated, block-sparse attention) could allow us to train AI models with much longer context. We are also excited to work with compiler researchers to make these optimization techniques easily programmable.

繁體中文 TRANSLATION

近期我們計畫與研究人員及工程師合作,讓 FlashAttention 更廣泛適用於不同裝置(例如 H100 GPU、AMD GPU),以及 FP8 等新資料型別。接下來的直接目標,是針對 H100 GPU 最佳化 FlashAttention-2,以運用其新硬體功能(TMA、第四代 Tensor Cores、fp8)。將 FlashAttention-2 的底層最佳化與局部、膨脹式、區塊稀疏注意力等高層演算法變更結合,可能讓我們訓練具有長得多上下文的 AI 模型。我們也期待與編譯器研究人員合作,讓這些最佳化技術更容易編寫。