RESEARCH NOTE 雙語閱讀 · 重點批註

FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision

Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri DaoarXiv preprint, arXiv:2407.086082024

arXiv:2407.08608v2

模型:gpt-6-sol · REASONING:high · DETAIL:high
輸入:40,299 tokens · 輸出:23,578 tokens(思考:3,977 tokens、其他輸出:19,601 tokens) · 總計:63,877 tokens

一句話掌握THE TAKEAWAY

FlashAttention-3 利用 Hopper GPU 的非同步資料搬移與矩陣運算,並處理 FP8 的資料布局和量化誤差,使精確注意力運算更快。

研究問題

如何在維持精確注意力計算的前提下,利用新一代 GPU 的非同步執行與低精度硬體,提高速度而不讓數值誤差失控?

動機

FlashAttention-2 已減少注意力運算的記憶體讀寫,但在 H100 上相對於最佳化矩陣乘法,硬體利用率仍低;softmax 的依賴關係及 FP8 的布局、離群值問題也阻礙直接加速。

方法

以 producer–consumer warp-specialization 和循環共享記憶體緩衝區重疊載入與計算;以跨區塊的兩階段管線及 pingpong 排程重疊 GEMM 和 softmax;FP8 路徑則加入核心內轉置、暫存器布局轉換、區塊量化與 incoherent processing。

主要貢獻

  • 將 TMA 載入和 Tensor Core 運算分配給不同 warps,建立非同步的 producer–consumer 管線。
  • 重排相鄰注意力區塊的運算,讓非同步 WGMMA 與 softmax 部分重疊。
  • 使 FP8 注意力符合 WGMMA 的布局限制,並用區塊量化與正交變換降低量化誤差。

主要結果

  • 在 H100 SXM5 的測試中,FP16 前向傳播相對 FlashAttention-2 快約 1.5–2.0 倍,反向傳播快約 1.5–1.75 倍。
  • FP16 最高約達 740 TFLOPs/s;FP8 前向傳播接近 1.2 PFLOPs/s,但相對 cuDNN 的優劣依維度、序列長度及遮罩條件而變。
  • 在論文模擬離群值的數值測試中,FP8 FlashAttention-3 的 RMSE 為 9.1e-3,per-tensor scaling 基線為 2.4e-2。

限制

  • 主要評估對象是 Hopper H100;作者尚未驗證其他硬體上的同等收益。
  • 兩階段管線增加暫存器需求;更深的三階段管線在文中的實作反而較慢。
  • 作者將 LLM 推論最佳化、FP8 persistent kernel,以及低精度注意力對大規模訓練的影響列為未來工作。
重要度
回到頂端 ↑

1 Introduction:研究缺口與貢獻

舊版 FlashAttention 主要從減少記憶體流量及增加平行度著手;本文指出,新硬體的非同步能力與 FP8 吞吐量並不能靠原有同步式演算法自然取得。

p. 2

原文 SOURCE

More fundamentally, FlashAttention-2’s algorithm adheres to a simplified synchronous model and makes no explicit use of asynchrony and low-precision in its design. Asynchrony is a result of hardware specialization to accelerate the most important operations in a ML workload: specific hardware units performing matrix multiplication (Tensor Cores) or memory loading (Tensor Memory Accelerator – TMA), separate from the rest of the CUDA cores performing logic, integer, and floating point computation. Low precision such as FP8 in Hopper and FP4 in Blackwell, continuing the trend of FP16 (Pascal in 2017) and BF16 (Ampere in 2020), is a proven technique to get double or quadruple throughput for the same power and chip area. We review the capabilities afforded by Hopper in these directions in § 2.2. The technical challenge is to redesign FlashAttention-2 to make use of these hardware features: asynchrony requires overlapping computation between matmul and softmax even though one depends on the output of the other, and low-precision requires care to minimize quantization error, especially in the case of outlier features in LLMs [20, 54].

繁體中文 TRANSLATION

更根本的是,FlashAttention-2 的演算法遵循簡化的同步模型,在設計上並未明確使用非同步能力與低精度。非同步能力來自硬體的專門分工,以加速機器學習工作負載中最重要的操作:由特定硬體單元執行矩陣乘法(Tensor Cores)或記憶體載入(Tensor Memory Accelerator,TMA),有別於執行邏輯、整數及浮點計算的其他 CUDA cores。Hopper 的 FP8 與 Blackwell 的 FP4 等低精度格式,延續 FP16(2017 年的 Pascal)及 BF16(2020 年的 Ampere)的趨勢,是在相同功耗和晶片面積下取得兩倍或四倍吞吐量的已驗證技術。我們在 § 2.2 回顧 Hopper 在這兩方面提供的能力。技術上的挑戰是重新設計 FlashAttention-2 以使用這些硬體特性:非同步執行要求矩陣乘法與 softmax 重疊,儘管後者依賴前者的輸出;低精度則必須謹慎控制量化誤差,尤其是在大型語言模型具有離群特徵的情況下 [20, 54]。

p. 2

原文 SOURCE

To this end, we propose FlashAttention-3, which contributes and synthesizes three new ideas to further improve performance on newer GPU architectures: 1. Producer-Consumer asynchrony: We define a warp-specialized software pipelining scheme that exploits the asynchronous execution of data movement and Tensor Cores by splitting producers and consumers of data into separate warps, thereby extending the algorithm’s ability to hide memory and instruction issue latencies. 2. Hiding softmax under asynchronous block-wise GEMMs: We overlap the comparatively low-throughput non-GEMM operations involved in softmax, such as floating point multiply-add and exponential, with the asynchronous WGMMA instructions for GEMM. As part of this, we rework the FlashAttention-2 algorithm to circumvent certain sequential dependencies between softmax and the GEMMs. For example, in the 2-stage version of our algorithm, while softmax executes on one block of the scores matrix, WGMMA executes in the asynchronous proxy to compute the next block. 3. Hardware-accelerated low-precision GEMM: We adapt the forward pass algorithm to allow for targeting the FP8 Tensor Cores for GEMM, nearly doubling the measured TFLOPs/s. This requires bridging the different layout conformance requirements of WGMMA in terms of how blocks of FP32 accumulator and FP8 operand matrices are assumed to be laid out in memory. We use the techniques of block quantization and incoherent processing to mitigate the loss of accuracy that results from moving to FP8 precision.

繁體中文 TRANSLATION

為此,我們提出 FlashAttention-3,結合三項新構想,進一步提升較新 GPU 架構上的效能: 1. 生產者—消費者非同步執行:我們定義一種 warp-specialized 軟體管線方案,將資料的生產者與消費者分配至不同 warps,利用資料搬移與 Tensor Cores 的非同步執行,進一步隱藏記憶體延遲及指令發出延遲。 2. 以非同步的區塊 GEMM 隱藏 softmax:我們讓 softmax 中吞吐量相對較低的非 GEMM 操作,例如浮點乘加與指數運算,和執行 GEMM 的非同步 WGMMA 指令重疊。為此,我們重新安排 FlashAttention-2 演算法,以繞開 softmax 與 GEMM 之間部分循序依賴。例如,在演算法的兩階段版本中,當 softmax 處理分數矩陣的一個區塊時,WGMMA 會在非同步執行路徑中計算下一區塊。 3. 硬體加速的低精度 GEMM:我們調整前向傳播演算法,使 GEMM 能使用 FP8 Tensor Cores,實測 TFLOPs/s 幾乎倍增。這需要處理 WGMMA 對 FP32 累加器區塊與 FP8 運算元矩陣區塊在記憶體中布局方式的不同要求。我們採用區塊量化與 incoherent processing,減輕改用 FP8 精度造成的準確度損失。

2 Background:計算目標與硬體模型

本文仍計算一般注意力輸出;設計空間來自 GPU 上彼此獨立的資料載入與矩陣乘法硬體,以及片上記憶體和暫存器的限制。

p. 2

原文 SOURCE

Let \(Q, K, V \in \mathbb{R}^{N\times d}\) be the query, key and value input sequences associated to a single head, where \(N\) is the sequence length and \(d\) is the head dimension. Then the attention output \(O\) is computed as: \[S = \alpha QK^\top \in \mathbb{R}^{N\times N},\quad P = \operatorname{softmax}(S) \in \mathbb{R}^{N\times N},\quad O = PV \in \mathbb{R}^{N\times d},\] where softmax is applied row-wise and one typically sets \(\alpha = 1/\sqrt{d}\) as the scaling factor. In practice, we subtract \(\operatorname{rowmax}(S)\) from \(S\) to prevent numerical instability with the exponential function. For multi-head attention (MHA), each head has its own set of query, key and value projections, and this computation parallelizes across multiple heads and batches to produce the full output tensor.

繁體中文 TRANSLATION

設 \(Q, K, V \in \mathbb{R}^{N\times d}\) 為單一注意力頭的查詢、鍵與值輸入序列,其中 \(N\) 是序列長度,\(d\) 是頭維度。注意力輸出 \(O\) 的計算方式為: \[S = \alpha QK^\top \in \mathbb{R}^{N\times N},\quad P = \operatorname{softmax}(S) \in \mathbb{R}^{N\times N},\quad O = PV \in \mathbb{R}^{N\times d},\] 其中 softmax 逐列套用,通常將 \(\alpha = 1/\sqrt{d}\) 設為縮放因子。實務上,為避免指數函數出現數值不穩定,會從 \(S\) 減去 \(\operatorname{rowmax}(S)\)。對多頭注意力(MHA),每個頭各有一組查詢、鍵和值的投影;此計算可跨多個頭和批次平行執行,產生完整輸出張量。

p. 3

原文 SOURCE

Asynchrony and warp-specialization: GPUs are throughput processors that rely on concurrency and asynchrony to hide memory and execution latencies. For async memory copy between GMEM and SMEM, Hopper has the Tensor Memory Accelerator (TMA) as a dedicated hardware unit [38, §7.29]. Furthermore, unlike prior architectures such as Ampere, the Tensor Core of Hopper, exposed via the warpgroup-wide WGMMA instruction [40, §9.7.14], is also asynchronous and can source its inputs directly from shared memory. Hardware support for asynchrony allows for warp-specialized kernels, where the warps of a CTA are divided into producer or consumer roles that only ever issue either data movement or computation. Generically, this improves the compiler’s ability to generate optimal instruction schedules [4]. In addition, Hopper supports the dynamic reallocation of registers between warpgroups via setmaxnreg [40, §9.7.17.1], so those warps doing MMAs can obtain a larger share of RMEM than those just issuing TMA (for which only a single thread is needed).

繁體中文 TRANSLATION

非同步執行與 warp-specialization:GPU 是依賴並行與非同步執行來隱藏記憶體及執行延遲的高吞吐量處理器。對於 GMEM 與 SMEM 之間的非同步記憶體複製,Hopper 有專用硬體單元 Tensor Memory Accelerator(TMA)[38, §7.29]。此外,與 Ampere 等先前架構不同,Hopper 的 Tensor Core 透過涵蓋整個 warpgroup 的 WGMMA 指令提供介面 [40, §9.7.14],也能非同步執行,且可直接從共享記憶體取得輸入。 硬體支援非同步執行,使 warp-specialized kernel 成為可能:將一個 CTA 的 warps 劃分為生產者或消費者,各自只發出資料搬移或計算指令。一般而言,這會改善編譯器產生最佳指令排程的能力 [4]。此外,Hopper 支援以 setmaxnreg 在 warpgroups 之間動態重新分配暫存器 [40, §9.7.17.1],因此執行矩陣乘加的 warps 可以取得比只發出 TMA 指令的 warps 更多 RMEM;後者僅需一個執行緒。

3.1 Producer–consumer 非同步與 pingpong 排程

每個 CTA 處理一塊查詢與對應輸出;生產者預先載入鍵和值,消費者執行注意力計算。兩個消費者 warpgroup 再交錯安排 GEMM 與 softmax。

p. 4

原文 SOURCE

Warp-specialization As with FlashAttention-2, the forward pass of FlashAttention-3 is embarrassingly parallel in the batch size, number of heads, and query sequence length. Thus, it will suffice to give a CTA-level view of the algorithm, which operates on a tile \(Q_i\) of the query matrix to compute the corresponding tile \(O_i\) of the output. To simplify the description, we first give the warp-specialization scheme with a circular SMEM buffer that does not have in addition the GEMM-softmax overlapping. Let \(d\) be the head dimension, \(N\) the sequence length, and fix a query block size \(B_r\) to divide \(Q\) into \(T_r = \lceil N/B_r\rceil\) blocks \(Q_1, .., Q_{T_r}\).

繁體中文 TRANSLATION

Warp-specialization:與 FlashAttention-2 相同,FlashAttention-3 的前向傳播可沿批次大小、頭數及查詢序列長度直接平行化。因此,只需從 CTA 層級描述演算法:它處理查詢矩陣的一個區塊 \(Q_i\),計算對應的輸出區塊 \(O_i\)。為簡化說明,我們先介紹採用環狀 SMEM 緩衝區的 warp-specialization 方案,此時尚未加入 GEMM 與 softmax 的重疊。設 \(d\) 為頭維度、\(N\) 為序列長度,並固定查詢區塊大小 \(B_r\),將 \(Q\) 分為 \(T_r = \lceil N/B_r\rceil\) 個區塊 \(Q_1, .., Q_{T_r}\)。

p. 5

原文 SOURCE

For our implementation of Algorithm 1 on Hopper, we use setmaxnreg for (de)allocations, TMA for loads of \(Q_i\) and \(\{K_j, V_j\}_{0\leq j<T_c}\), and WGMMA to execute the GEMMs in the consumer mainloop, with the SS or RS prefix indicating whether the first operand is sourced from shared memory or register file. For interpreting the execution flow of Algorithm 1, note that issuing TMA loads does not stall on the completion of other loads due to asynchrony. Moreover, in the producer mainloop, no waits will be issued for the first \(s\) iterations as the buffer gets filled.

繁體中文 TRANSLATION

在 Hopper 上實作 Algorithm 1 時,我們以 setmaxnreg 進行暫存器的配置與釋放,以 TMA 載入 \(Q_i\) 及 \(\{K_j, V_j\}_{0\leq j<T_c}\),並以 WGMMA 執行消費者主迴圈中的 GEMM;其中 SS 或 RS 前綴指出第一個運算元來自共享記憶體或暫存器檔案。解讀 Algorithm 1 的執行流程時,須注意非同步特性使發出 TMA 載入指令不必停下來等待其他載入完成。此外,生產者主迴圈在緩衝區填滿之前的前 \(s\) 次迭代不會發出等待指令。

p. 5

原文 SOURCE

Since the exponential is performed by a separate hardware unit (the multi-function unit), ideally we’d want the exponential calculation to be scheduled when the Tensor Cores are performing the matmul. To do so, we use synchronization barriers (bar.sync instructions) to force the GEMMs (GEMM1 – PV of one iteration, and GEMM0 – \(QK^\top\) of the next iteration) of warpgroup 1 to be scheduled before the GEMMs of warpgroup 2. As a result, the softmax of warpgroup 1 will be scheduled while warpgroup 2 is performing its GEMMs. Then the roles swap, with warpgroup 2 doing softmax while warpgroup 1 doing GEMMs (hence, “pingpong” scheduling). This is illustrated in Fig. 1. Though in practice the pingpong scheduling is not as clean as depicted in the figure, we generally find this to improve performance (e.g., from 570 TFLOPS to 620-640 TFLOPS for FP16 forward with head dimension 128 and sequence length 8192).

繁體中文 TRANSLATION

由於指數運算由另一個硬體單元(multi-function unit)執行,理想上我們希望在 Tensor Cores 執行矩陣乘法時安排指數計算。為此,我們使用同步屏障(bar.sync 指令),強制 warpgroup 1 的 GEMM——某次迭代的 GEMM1,即 \(PV\),以及下一次迭代的 GEMM0,即 \(QK^\top\)——排在 warpgroup 2 的 GEMM 之前。因此,當 warpgroup 2 執行其 GEMM 時,warpgroup 1 會執行 softmax。接著兩組交換角色:warpgroup 2 執行 softmax,而 warpgroup 1 執行 GEMM,因此稱為「pingpong」排程。Fig. 1 說明此安排。雖然實務上的 pingpong 排程不像圖中那麼整齊,我們通常發現它能改善效能;例如,在頭維度為 128、序列長度為 8192 的 FP16 前向傳播中,從 570 TFLOPS 提升至 620–640 TFLOPS。

Figure 1原始 PDF 第 6 頁

Pingpong scheduling for 2 warpgroups to overlap softmax and GEMMs: the softmax of one warpgroup should be scheduled when the GEMMs of another warpgroup are running. The same color denotes the same iteration.兩個 warpgroups 的 pingpong 排程,用以重疊 softmax 與 GEMM:當一個 warpgroup 執行 softmax 時,應讓另一個 warpgroup 執行 GEMM。相同顏色表示同一次迭代。

這張圖在說什麼

圖示兩個 warpgroup 的 GEMM 與 softmax 如何沿時間交錯。它呈現作者希望達成的排程關係,而非逐指令的實際執行紀錄。

怎麼看

先沿水平時間軸看上、下兩列 warpgroup;GEMM0 是分數矩陣乘法,GEMM1 是與值矩陣相乘。再看同一時間區間內一列的 softmax 是否對上另一列的 GEMM;顏色用來辨認迭代,虛線協助辨認排程階段,沒有縱軸效能數值。

3.2 同一 warpgroup 內的 GEMM–softmax 管線

作者跨相鄰區塊重排運算,以額外暫存器保存中間結果,使原本受資料依賴限制的 softmax 與非同步 WGMMA 得以部分重疊。

p. 6

原文 SOURCE

In the attention algorithm, operations within the inner loop (main loop) have sequential dependencies that impede parallelization within a single iteration. For example, (local) softmax (lines 18 to 19) relies on the output \(S_i^{(j)}\) of the first GEMM, while the second GEMM takes its result \(\widetilde{P}_i^{(j)}\) as an operand. Indeed, the wait statements in lines 17 and 21 of Algorithm 1 serialize the execution of softmax and GEMMs. However, we can break these dependencies by pipelining across iterations through additional buffers in registers. Pursuing this idea, we propose the following two-stage GEMM-softmax pipelining algorithm:

繁體中文 TRANSLATION

在注意力演算法中,內部迴圈(主迴圈)的操作有循序依賴,阻礙單次迭代內的平行化。例如,(局部)softmax(第 18 至 19 行)依賴第一次 GEMM 的輸出 \(S_i^{(j)}\),而第二次 GEMM 則以其結果 \(\widetilde{P}_i^{(j)}\) 為運算元。確實,Algorithm 1 第 17 與 21 行的等待敘述,使 softmax 和 GEMM 依序執行。不過,藉由暫存器中的額外緩衝區,在不同迭代之間建立管線,我們可以突破這些依賴。循此構想,我們提出以下兩階段 GEMM–softmax 管線演算法:

Figure 2原始 PDF 第 6 頁

2-stage WGMMA-softmax pipelining兩階段 WGMMA–softmax 管線

這張圖在說什麼

圖中以不同迭代的 WGMMA0、softmax、WGMMA1 表示跨迭代重疊。它支撐作者的主張:將下一區塊的分數計算與相鄰區塊的 softmax/輸出計算交錯,而非完全序列化。

怎麼看

先沿水平時間軸讀三列:WGMMA0 計算分數、Softmax 正規化分數、WGMMA1 計算輸出;方塊內數字是迭代編號。比較不同列在同一時段是否出現不同編號的工作;圖是排程示意,並非測得的時間或速度曲線。

p. 7

原文 SOURCE

Algorithm 2 functions as a replacement for the consumer path of Algorithm 1 to comprise the complete FlashAttention-3 algorithm for FP16 precision. At a high-level, we use WGMMA as a metonym for asynchronous GEMM. Within the mainloop (lines 8 to 16), the second WGMMA operation of iteration \(j\) (line 11) is overlapped with softmax operations from iteration \(j + 1\) (line 13).

繁體中文 TRANSLATION

Algorithm 2 取代 Algorithm 1 的消費者路徑,兩者共同構成 FP16 精度下完整的 FlashAttention-3 演算法。在高層次的描述中,我們以 WGMMA 代指非同步 GEMM。在主迴圈(第 8 至 16 行)中,迭代 \(j\) 的第二個 WGMMA 操作(第 11 行)與迭代 \(j + 1\) 的 softmax 操作(第 13 行)重疊。

3.3 FP8:布局相容性與誤差控制

FP8 WGMMA 的運算元布局要求使第二次 GEMM 需要轉置值矩陣區塊,且兩次 GEMM 之間需要暫存器重排。區塊量化與隨機正交變換則處理精度下降。

p. 8

原文 SOURCE

Instead, for FP8 FlashAttention-3 we opt for option (2). For the in-kernel transpose, we take advantage of the LDSM (ldmatrix) and STSM (stmatrix) instructions, which involve a warp of threads collectively loading SMEM to RMEM and storing RMEM to SMEM at a granularity of 128 bytes. The LDSM/STSM instructions are both register efficient, allowing us to execute them in the producer warpgroup, and capable of transposing layouts when doing memory copy. Moreover, after the first iteration we can arrange for the transpose of the next V tile to be executed in the shadow of the two WGMMAs that involve the preceding V and current K tile.

繁體中文 TRANSLATION

因此,對 FP8 FlashAttention-3,我們選擇方案(2)。進行核心內轉置時,我們利用 LDSM(ldmatrix)與 STSM(stmatrix)指令,讓一個 warp 的執行緒共同以 128 位元組為粒度,將資料由 SMEM 載入 RMEM,再由 RMEM 存回 SMEM。LDSM/STSM 指令既節省暫存器,使我們能在生產者 warpgroup 中執行,也能在記憶體複製時轉換布局。此外,第一次迭代之後,我們可以安排下一個 \(V\) 區塊的轉置,在涉及前一個 \(V\) 區塊及目前 \(K\) 區塊的兩個 WGMMA 執行期間完成。

p. 8

原文 SOURCE

Second, we observe that unlike with FP16, the memory layout of the FP32 accumulator of an FP8 WGMMA is different from that assumed for its operand A when held in registers. We depict fragments of these two layouts in Fig. 3 and Fig. 4, where the entries are held in registers per thread in the listed order. By using byte permute instructions, we can then transform the first WGMMA’s accumulator into a format suitable for the second WGMMA, and compatibly with the layout of the V tile produced by the in-kernel transpose. Specifically, with reference to Fig. 3, we change the order in sequence to \[\{d0\ d1\ d4\ d5\ d2\ d3\ d6\ d7\},\] and this register permutation is then replicated over every 8 bytes. In terms of the logical shape of the P tile, this manuever permutes its columns (e.g., columns 0189 now become the first four columns). For WGMMA to then compute the correct output tile, we can correspondingly arrange for the in-kernel transpose to write out a matching row permutation of the V tile.

繁體中文 TRANSLATION

其次,我們觀察到,與 FP16 不同,FP8 WGMMA 的 FP32 累加器在記憶體中的布局,不同於其運算元 A 位於暫存器時預期的布局。我們在 Fig. 3 和 Fig. 4 描繪這兩種布局的局部;圖中的項目按所列順序存於各執行緒的暫存器。藉由位元組置換指令,我們可以把第一次 WGMMA 的累加器轉成適合第二次 WGMMA 的格式,並與核心內轉置產生的 \(V\) 區塊布局相容。具體而言,參照 Fig. 3,我們將順序改為 \[\{d0\ d1\ d4\ d5\ d2\ d3\ d6\ d7\},\] 並對每個 8 位元組重複這項暫存器置換。就 \(P\) 區塊的邏輯形狀而言,此操作置換了它的欄(例如,欄 0189 現在成為前四欄)。為讓 WGMMA 隨後算出正確的輸出區塊,我們可相應安排核心內轉置,使其輸出具有匹配之列置換的 \(V\) 區塊。

Figure 3原始 PDF 第 8 頁

FP32 accumulator register WGMMA layout – rows 0 and 8, threads 0-3, entries 0-7.FP32 累加器的 WGMMA 暫存器布局——第 0 與第 8 列、執行緒 0–3、項目 0–7。

這張圖在說什麼

圖中列出第一次 WGMMA 的 FP32 累加結果在部分執行緒暫存器中的配置。它是辨認後續位元組置換來源順序的對照圖。

怎麼看

先看每格的 T0–T3,辨認持有資料的執行緒;再看大括號內的 d0–d7,辨認各暫存器的項目順序。只展示指定列與執行緒的局部布局,不表示完整矩陣,也沒有速度座標軸。

Figure 4原始 PDF 第 8 頁

FP8 operand A register WGMMA layout – rows 0 and 8, threads 0-3, entries 0-7.FP8 運算元 A 的 WGMMA 暫存器布局——第 0 與第 8 列、執行緒 0–3、項目 0–7。

這張圖在說什麼

圖中顯示第二次 WGMMA 對 FP8 運算元 A 的部分暫存器布局要求。與 Figure 3 對照,可看出累加器輸出不能不經重排就直接沿用。

怎麼看

依序比較每格的執行緒 T0–T3 和項目 a0–a7,再與 Figure 3 相同位置的項目順序對照。此圖說明資料配置要求,而非數值精度或吞吐量。

p. 8

原文 SOURCE

Accuracy: block quantization and incoherent processing. With FP8 (e4m3) format, one only uses 3 bits to store the mantissa and 4 bits for the exponent. This results in higher numerical error than FP16/BF16. Moreover, large models typically have outlier values [20, 54] that are much larger in magnitude than most other values, making quantization difficult. One typically use per-tensor scaling [37] by keeping one scalar per tensor (e.g., one for Q, for K, and for V). To reduce the numerical error of attention in FP8, we employ two techniques: 1. Block quantization: we keep one scalar per block, so that for each of Q, K, V we split the tensor into blocks of size \(B_r \times d\) or \(B_c \times d\) and quantize them separately. This quantization can be fused with an operation right before attention (e.g., rotary embedding) with no additional slow down (since rotary embedding is memory-bandwidth bound). As the FlashAttention-3 algorithm naturally operates on blocks, we can scale each block of S to account for this block quantization at no computation cost.

繁體中文 TRANSLATION

準確度:區塊量化與 incoherent processing。FP8(e4m3)格式僅以 3 個位元儲存尾數、4 個位元儲存指數,因此數值誤差高於 FP16/BF16。此外,大型模型通常含有大小遠超過其他數值的離群值 [20, 54],使量化更加困難。一般的 per-tensor scaling [37] 為每個張量保存一個縮放純量,例如 \(Q\)、\(K\)、\(V\) 各一個。為降低 FP8 注意力的數值誤差,我們採用兩項技術: 1. 區塊量化:每個區塊保存一個縮放純量;對 \(Q\)、\(K\)、\(V\) 各張量,將其切成大小為 \(B_r \times d\) 或 \(B_c \times d\) 的區塊,分別量化。量化可以與注意力之前的操作(例如 rotary embedding)融合,而不增加執行時間,因為 rotary embedding 受記憶體頻寬限制。由於 FlashAttention-3 原本就以區塊運作,我們可以縮放 \(S\) 的各區塊,以補償區塊量化,而不增加計算成本。

p. 9

原文 SOURCE

2. Incoherent processing: to even out outliers, we multiply Q and K with a random orthogonal matrix M before quantizing to FP8. Since M is orthogonal, \(MM^\top = I\) and so \((QM)(KM)^\top = QK^\top\), i.e., multiplying both Q and K with M does not change the attention output. This serves to “spread out” the outliers since each entry of QM or KM is a random sum of entries of Q or K, thus reducing quantization error. In practice, we follow Chee et al. [9] and Tseng et al. [58] and choose M to be the product of random diagonal matrices of ±1 and a Hadamard matrix, which can be multiplied in \(O(d \log d)\) instead of \(O(d^2)\), and can also be fused with the rotary embedding at no extra computation cost.

繁體中文 TRANSLATION

2. Incoherent processing:為使離群值的影響分散,我們在量化為 FP8 之前,將 \(Q\) 和 \(K\) 乘上一個隨機正交矩陣 \(M\)。由於 \(M\) 正交,\(MM^\top = I\),因此 \((QM)(KM)^\top = QK^\top\);也就是說,讓 \(Q\) 和 \(K\) 同乘 \(M\) 不會改變注意力輸出。這會將離群值「攤散」,因為 \(QM\) 或 \(KM\) 的每個元素都是 \(Q\) 或 \(K\) 中多個元素的隨機加總,因而降低量化誤差。實務上,我們依循 Chee 等人 [9] 與 Tseng 等人 [58],選用由元素為 ±1 的隨機對角矩陣與 Hadamard 矩陣相乘所構成的 \(M\);其乘法可用 \(O(d \log d)\) 而非 \(O(d^2)\) 的成本執行,也能與 rotary embedding 融合而不增加額外計算成本。

4.1 效能基準測試

H100 上的前向與反向測試支持 FP16 加速,但 FP8 與 cuDNN 的比較須依頭維度、遮罩及序列長度分開閱讀。

p. 9

原文 SOURCE

We measure the runtime of different attention methods on an H100 80GB SXM5 GPU for different settings (without / with causal mask, head dimension 64 or 128) for FP16 inputs. We report the results in Fig. 5 and Fig. 6, showing that FlashAttention-3 is around 1.5-2.0× faster than FlashAttention-2 in the forward pass and 1.5-1.75× faster in the backward pass. Compared to a standard attention implementation, FlashAttention-3 can be up to 3-16× faster. For medium and long sequences (1k and above), FlashAttention-3 even surpasses the speed of a vendor’s library (cuDNN – closed source) that has been optimized for H100 GPUs.

繁體中文 TRANSLATION

我們在 H100 80GB SXM5 GPU 上,以 FP16 輸入測量不同注意力方法在不同設定下的執行時間,包括有無因果遮罩及頭維度 64 或 128。Fig. 5 和 Fig. 6 呈現結果:FlashAttention-3 的前向傳播約比 FlashAttention-2 快 1.5–2.0 倍,反向傳播快 1.5–1.75 倍。與標準注意力實作相比,FlashAttention-3 最多可快 3–16 倍。對中長序列(1k 以上),FlashAttention-3 甚至快於已針對 H100 GPU 最佳化的廠商函式庫 cuDNN(封閉原始碼)。

Figure 5原始 PDF 第 10 頁

Attention forward speed (FP16/BF16) on H100 GPUH100 GPU 上的注意力前向傳播速度(FP16/BF16)

這張圖在說什麼

六個子圖比較不同頭維度及有無因果遮罩時的前向吞吐量。圖中 FlashAttention-3 在較長序列通常明顯高於 FlashAttention-2;與 cuDNN 的差距則會隨設定改變。

怎麼看

先按子圖標題辨認頭維度 64、128、256,再看左欄無因果遮罩、右欄有因果遮罩。橫軸為序列長度,縱軸為 TFLOPs/s,越高越好;於同一子圖和同一序列長度比較各色柱,並留意標成 OOM 的基線沒有可比較數值。

Figure 6原始 PDF 第 11 頁

Attention backward speed (FP16/BF16) on H100 GPUH100 GPU 上的注意力反向傳播速度(FP16/BF16)

這張圖在說什麼

兩個子圖給出無因果遮罩時、頭維度 64 與 128 的反向吞吐量。它支持反向傳播也有加速,但不能據此推論未展示的反向遮罩設定。

怎麼看

先看左圖頭維度 64、右圖頭維度 128;橫軸是序列長度,縱軸 TFLOPs/s 越高越好。在相同長度比較 FlashAttention-3、FlashAttention-2、cuDNN 與標準注意力,並將 OOM 視為未能取得該基線數值,而非速度為零。

p. 9

原文 SOURCE

We also measure the runtime for FP8 for the forward pass under similar settings. We report the results for headdim 256 in Fig. 7 and give the full results in Appendix C.2.

繁體中文 TRANSLATION

我們也在類似設定下測量 FP8 前向傳播的執行時間。Fig. 7 報告頭維度 256 的結果,完整結果則列於 Appendix C.2。

Figure 7原始 PDF 第 11 頁

Attention forward speed (FP8) on H100 GPUH100 GPU 上的注意力前向傳播速度(FP8)

這張圖在說什麼

圖中比較頭維度 256、FP8 前向傳播的 Triton、cuDNN 與 FlashAttention-3。無因果遮罩的長序列下 FlashAttention-3 接近 1.2 PFLOPs/s;有因果遮罩時,cuDNN 在圖示的多個較長序列設定較快。

怎麼看

先分辨左圖無因果遮罩、右圖有因果遮罩;兩圖頭維度均為 256。橫軸是序列長度,縱軸 TFLOPs/s 越高越好;在同一長度比較綠色 Triton、紅色 cuDNN、紫色 FlashAttention-3,不要把左圖的領先直接套到右圖。

4.2 消融實驗

固定一組非因果 FP16 參數後,移除 GEMM–softmax 管線或 warp-specialization 都會降低測得吞吐量。

p. 9

原文 SOURCE

We ablate both the 2-stage WGMMA-softmax pipelining and warp-specialization for non-causal FP16 FlashAttention-3 with fixed parameters {batch,seqlen, nheads, hdim} = {4, 8448, 16, 128}. The result in Table 2 confirms that our algorithmic improvements (asynchrony with warp-specialization and overlapping between GEMM and softmax) lead to significant speedup, from 570 to 661 TFLOPs.

繁體中文 TRANSLATION

我們在非因果 FP16 FlashAttention-3 上,以固定參數 {batch,seqlen, nheads, hdim} = {4, 8448, 16, 128},分別消融兩階段 WGMMA–softmax 管線與 warp-specialization。Table 2 的結果證實,我們的演算法改進——結合 warp-specialization 的非同步執行,以及 GEMM 與 softmax 的重疊——帶來顯著加速,從 570 提升至 661 TFLOPs。

Table 2原始 PDF 第 11 頁

Pipelining ablation measurements管線設計消融測量結果
ConfigurationTimeTFLOPs/s
FlashAttention-33.538 ms661
No GEMM-Softmax Pipelining, Warp-Specialization4.021 ms582
GEMM-Softmax Pipelining, No Warp-Specialization4.105 ms570

表格註記:測試條件見引文:非因果 FP16,{batch,seqlen, nheads, hdim} = {4, 8448, 16, 128}。表格未列兩項機制同時移除的配置。

這張表在說什麼

表格在固定非因果 FP16 工作負載下,比較完整 FlashAttention-3 與各移除一項機制後的時間及吞吐量。完整版本時間最短、吞吐量最高。

怎麼看

先看 Configuration 確認移除的是哪一機制,再比較 Time 與 TFLOPs/s;Time 越低越好,TFLOPs/s 越高越好。三列使用相同的 {batch,seqlen, nheads, hdim} 設定,不能直接外推到其他維度或遮罩條件。

4.3 數值誤差驗證

作者對模擬離群值的輸入,以 FP64 作為參考比較 RMSE。FP16 FlashAttention-3 與前代誤差相同;FP8 改進版則優於 per-tensor scaling 基線。

p. 10

原文 SOURCE

As there has been interest in the numerical error [21] of FlashAttention, we compare FlashAttention-2, FlashAttention-3, and a standard implementation of attention against a reference implementation in FP64. To simulate outlier features and activations in LLMs [20, 54], we generate the entries of Q, K, V with the following distribution: \[\mathcal{N}(0,1) + \mathcal{N}(0,100)\cdot\operatorname{Bernoulli}(0.001).\] That is, each entry is normally distributed with zero mean and standard deviation 1, but for 0.1% of entries we add an independent term that’s normally distributed with standard deviation 10. We then measure the root mean squared error (RMSE) in Table 3. In FP16, both FlashAttention-2 and FlashAttention-3 achieves 1.7× lower RMSE compared to the standard implementation since intermediate results (softmax) are kept in FP32. The baseline attention in FP8 uses per-tensor scaling, with matmul accumulator in FP32 and intermediate softmax results kept in FP16. Thanks to block quantization and incoherent processing, FlashAttention-3 in FP8 is 2.6× more accurate than this baseline.

繁體中文 TRANSLATION

由於 FlashAttention 的數值誤差受到關注 [21],我們將 FlashAttention-2、FlashAttention-3 與標準注意力實作,和 FP64 的參考實作比較。為模擬大型語言模型中的離群特徵與活化值 [20, 54],我們用以下分布產生 \(Q\)、\(K\)、\(V\) 的元素: \[\mathcal{N}(0,1) + \mathcal{N}(0,100)\cdot\operatorname{Bernoulli}(0.001).\] 也就是說,每個元素原本服從平均數為零、標準差為 1 的常態分布,但有 0.1% 的元素會再加上一個標準差為 10 的獨立常態分布項。接著,我們在 Table 3 測量均方根誤差(RMSE)。在 FP16 下,FlashAttention-2 和 FlashAttention-3 的 RMSE 都比標準實作低 1.7 倍,因為中間結果(softmax)維持在 FP32。FP8 注意力基線使用 per-tensor scaling,矩陣乘法累加器採 FP32,而 softmax 的中間結果保留在 FP16。藉由區塊量化與 incoherent processing,FP8 FlashAttention-3 的準確度比此基線高 2.6 倍。

Table 3原始 PDF 第 12 頁

Numerical error comparisons in FP16 and FP8 (e4m3).FP16 與 FP8(e4m3)的數值誤差比較。
MetricBaseline FP16FlashAttention-2 FP16FlashAttention-3 FP16Baseline FP8FlashAttention-3 FP8No block quantNo incoherent processing
RMSE3.2e-41.9e-41.9e-42.4e-29.1e-39.3e-32.4e-2

表格註記:原表依序分為上下兩組,每組各有一列 Method 與一列 RMSE;此處將兩組方法欄並列展平為一列數值,保留各組內原始欄序。比較條件為文中模擬離群值的輸入及 FP64 參考實作。

這張表在說什麼

上半組顯示 FP16 的兩個 FlashAttention 版本有相同 RMSE,且均低於標準基線;下半組顯示完整 FP8 方法低於 per-tensor scaling 基線。兩個 FP8 消融欄有助於辨認各技巧在此測試下的影響。

怎麼看

先分別讀 FP16 三個方法與 FP8 四個方法,不要跨精度把 RMSE 當作同一量化條件的消融。所有數值都是相對 FP64 參考實作的 RMSE,越低越好;讀 FP8 消融時,比較完整方法、No block quant、No incoherent processing 和 Baseline FP8。

5 Discussion, Limitations, Conclusion:結論與邊界

作者總結 H100 上的速度與誤差收益,同時明確指出推論、FP8 kernel 設計和大規模訓練尚待研究。

p. 12

原文 SOURCE

With FlashAttention-3, we have demonstrated that new programming techniques and hardware features such as asynchrony and low-precision can have a dramatic impact on the efficiency and accuracy of attention. We are able to speed up attention by 1.5-2.0× times compared to FlashAttention-2, and reduce FP8 numerical error by 2.6× compared to standard per-tensor quantization. Some limitations of our work that we hope to address in the future include: optimizing for LLM inference, integrating a persistent kernel design into the FP8 kernel, and understanding the effects of low-precision attention in large-scale training. Though we have focused on Hopper GPUs in this work, we expect that the techniques developed here will apply to other hardware accelerators. We hope that a faster and more accurate primitive such as attention will unlock new applications in long-context tasks.

繁體中文 TRANSLATION

透過 FlashAttention-3,我們展示非同步執行與低精度等新程式設計技術及硬體特性,能顯著影響注意力運算的效率與準確度。相對於 FlashAttention-2,我們使注意力運算加速 1.5–2.0 倍;相對於標準 per-tensor quantization,則使 FP8 數值誤差降低 2.6 倍。我們希望未來處理的研究限制包括:最佳化大型語言模型推論、將 persistent kernel 設計整合至 FP8 kernel,以及了解低精度注意力對大規模訓練的影響。雖然本研究聚焦於 Hopper GPU,我們預期所開發的技術也適用於其他硬體加速器。我們希望更快且更準確的注意力這類基礎運算,能促成長上下文任務的新應用。