電路板上發光的 AI 晶片特寫

← INSIGHTS & PERSPECTIVES | 機器學習

FlashAttention 介紹:更快的注意力機制如何省記憶體又加速 Transformer 訓練

FlashAttention 與 FlashAttention-2 官方實作介紹:IO 感知的精確注意力演算法,透過切片重算與硬體特性優化,讓 Transformer 訓練更快、記憶體用量從 O(N²) 降到線性,並整理安裝需求與 flash_attn 使用範例。

我在研究 Transformer 模型加速時接觸到 FlashAttention,它是一種快速且記憶體高效的精確注意力機制,設計時考慮了 IO(輸入輸出)的特性。這篇筆記整理 FlashAttention 與 FlashAttention-2 的官方資訊、核心特點、原理概念,以及實際安裝與使用的步驟。

FlashAttention 的官方資源在哪裡?

官方資訊如下:

此存儲庫提供了以下論文中 FlashAttention 和 FlashAttention-2 的官方實現,可讓我們在建模時有更快的注意力、更好的並行度和工作分區。下面為一張概念示意圖:

FlashAttention 概念示意圖

官方所做的效能提升試驗結果如下:

FlashAttention 效能提升試驗結果

甚麼是 Flash Attention?

Flash Attention 是一種注意力算法,旨在提高基於 Transformer 的模型的效率,使其能夠處理更長的序列長度並更快地進行訓練和推理。它通過減少計算量和內存使用來實現這一點。Flash Attention 是一種快速且內存高效的精確注意力機制,其設計考慮了 IO(輸入輸出)的特性。

這項技術的關鍵點有哪些?

快速 (Fast)

  • 訓練 BERT-large(序列長度 512)比 MLPerf 1.1 中的訓練速度記錄快 15%。
  • 訓練 GPT-2(序列長度 1K)比 HuggingFace 和 Megatron-LM 的基準實現快 3 倍。
  • 在 long-range arena(序列長度 1K-4K)中,比基準速度快 2.4 倍。

高效內存使用 (Memory-efficient)

  • 傳統的注意力機制內存訪問量是 O(N²),而 Flash Attention 的內存訪問量是亞二次方/線性的。

精確 (Exact)

  • 這不是近似算法(例如稀疏或低秩矩陣方法),其結果與原始方法完全相同。

IO感知 (IO-aware)

  • 與原始的注意力計算方法相比,Flash Attention 考慮了硬件(特別是 GPU)的特性,而不是將其當作黑盒來處理。

如何使用 Flash Attention 實現加速?

可以通過以下兩種方式來實現:

  • 切片和重新計算:Flash Attention 將序列分成較小的塊,並在每個塊上計算注意力。這可以減少計算量,因為每個塊的注意力矩陣都小得多。此外,Flash Attention 還會重新利用中間計算結果,以進一步減少計算量。
  • 稀疏表示:Flash Attention 使用稀疏表示來表示注意力矩陣。這意味著只存儲非零元素,從而減少內存使用量。

怎麼安裝 Flash Attention?

系統要求:

  • CUDA 11.6 及更高版本。
  • PyTorch 1.12 及更高版本。
  • Linux 系統。此功能有可能於 v2.3.2 版本之後開始支持 Windows,但 Windows 編譯仍然需要更多的測試。

我們推薦 Nvidia 的 Pytorch 容器,它具有安裝 FlashAttention 所需的所有工具。

在使用 Flash Attention 前要先安裝:

  1. PyTorch。
  2. 安裝 `pip install packaging`
  3. 確保已安裝並且 `ninja` 工作正常(例如,`ninja --version` 然後 `echo $?` 應返回退出代碼 0)。如果不是(有時 `ninja --version` 然後 `echo $?` 返回非零退出代碼),請卸載然後重新安裝 `ninja`(`pip uninstall -y ninja && pip install ninja`)。如果沒有 `ninja`,編譯可能需要很長時間(2 小時),因為它不使用多個 CPU 內核。`ninja` 在 3 核機器上編譯需要 5-64 分鐘。
  4. 然後:

```bash

pip install flash-attn --no-build-isolation

```

如果您的電腦的 RAM 小於 96GB 且 CPU 內核眾多,`ninja` 則可能會運行過多的並行編譯作業,從而耗盡 RAM 量。要限制並行編譯作業的數量,可以設置環境變數 `MAX_JOBS`:

```bash

MAX_JOBS=4 pip install flash-attn --no-build-isolation

```

flash_attn 的程式碼怎麼寫?

使用範例:

```python

from flash_attn import flash_attn_qkvpacked_func, flash_attn_func

flash_attn_qkvpacked_func(qkv, dropout_p=0.0, softmax_scale=None, causal=False,

window_size=(-1, -1), alibi_slopes=None, deterministic=False):

"""dropout_p should be set to 0.0 during evaluation

If Q, K, V are already stacked into 1 tensor, this function will be faster than

calling flash_attn_func on Q, K, V since the backward pass avoids explicit concatenation

of the gradients of Q, K, V.

If window_size != (-1, -1), implements sliding window local attention. Query at position i

will only attend to keys between [i - window_size[0], i + window_size[1]] inclusive.

Arguments:

qkv: (batch_size, seqlen, 3, nheads, headdim)

dropout_p: float. Dropout probability.

softmax_scale: float. The scaling of QK^T before applying softmax.

Default to 1 / sqrt(headdim).

causal: bool. Whether to apply causal attention mask (e.g., for auto-regressive modeling).

window_size: (left, right). If not (-1, -1), implements sliding window local attention.

alibi_slopes: (nheads,) or (batch_size, nheads), fp32. A bias of (-alibi_slope * |i - j|) is added to

the attention score of query i and key j.

deterministic: bool. Whether to use the deterministic implementation of the backward pass,

which is slightly slower and uses more memory. The forward pass is always deterministic.

Return:

out: (batch_size, seqlen, nheads, headdim).

"""

```

延伸閱讀

常見問題

QFlashAttention 是近似算法嗎?

不是。FlashAttention 是精確(exact)的注意力機制,其計算結果與原始注意力方法完全相同,不像稀疏或低秩矩陣等近似方法會犧牲精度。

QFlashAttention 為什麼能節省記憶體?

傳統注意力機制的內存訪問量是 O(N²),而 FlashAttention 透過分塊計算(切片和重新計算)把注意力矩陣變小,內存訪問量降到亞二次方/線性。

Q安裝 flash-attn 需要什麼環境?

需要 CUDA 11.6 以上、PyTorch 1.12 以上,且目前主要支援 Linux。官方建議使用 Nvidia 的 PyTorch 容器,裡面已包含安裝所需工具。

Q編譯 flash-attn 卡很久怎麼辦?

先確認 `ninja` 已正確安裝(`ninja --version; echo $?` 回傳 0),沒有多核編譯可能要花上 2 小時。若 RAM 小於 96GB,記得用 `MAX_JOBS` 限制並行編譯數量,例如 `MAX_JOBS=4 pip install flash-attn --no-build-isolation`。

Q什麼時候該用 flash_attn_qkvpacked_func 而不是 flash_attn_func?

當 Q、K、V 已經堆疊成單一張量時,用 `flash_attn_qkvpacked_func` 會更快,因為反向傳播可以避免對 Q、K、V 的梯度做明確的串接。

參考資料

最後更新

2026-08-28(原文發布於 2024-07-24,本文保留原始筆記內容並補上 GEO 結構。)

關於作者 {#author}

Claire Chang | 企業 AI 導入與流程轉型顧問。專注於 AI Agent 架構設計、ERP 系統整合與企業 AI 治理。

首次發布:2024-07-24