FlashAttention-2/3
3 秒看懂
FlashAttention 是一種 IO 感知的精確注意力演算法——它不近似、不丟精度,而是通過分塊(tiling)+ 線上 Softmax 技巧,把注意力計算的視訊記憶體佔用從 O(N²) 降到 O(N),同時大幅減少對高頻寬視訊記憶體(HBM)的讀寫次數,從而在現代 GPU 上獲得 2–4× 的實際吞吐提升。FlashAttention-2(2023)在 A100 上實現了約 72% 的 FP16 Tensor Core 峰值利用率(約 230 TFLOPS,論文報告值);FlashAttention-3(2024)則針對 Hopper(H100)架構,利用非同步執行和 FP8 低精度進一步逼近硬體極限。
3 分鐘產業解釋
為什麼注意力機制是個瓶頸?
Transformer 的核心操作——自注意力(Self-Attention)——需要計算 Q·K^T 得到一個 N×N 的注意力矩陣(N 為序列長度)。當上下文視窗從 2K 擴充套件到 128K 甚至更長時,這個矩陣從幾百 MB 暴漲到數百 GB,不僅 視訊記憶體裝不下,而且頻繁的 HBM 讀寫(而非計算本身)成為真正的效能瓶頸。標準注意力實現在 A100 上通常只能利用 ~5–25% 的 Tensor Core 算力,大量時間花在等資料搬運。
FlashAttention 的商業意義
- 使能長上下文:GPT-4 的 128K、Claude 的 200K 視窗,底層訓練和推論幾乎都依賴 FlashAttention 系列。
- 降低訓練成本:同樣的 GPU 叢集,訓練吞吐提升 2×+ 意味著數百萬美元的節省。
- 減少視訊記憶體需求:不儲存完整注意力矩陣,使單卡可處理更長序列或更大 batch。
- 行業標準化:PyTorch 2.0+ 已將 FlashAttention 整合為
scaled_dot_product_attention的預設後端之一;Hugging Face Transformers 預設啟用。
誰是核心受益者?
| 角色 | 受益邏輯 |
|---|---|
| NVIDIA | 自家硬體利用率被拉滿,進一步鞏固 CUDA 生態 |
| 所有 LLM 訓練/推論廠商 | 同等算力預算下訓練更快、上下文更長 |
| HBM 供應商(SK海力士/三星/美光) | 長上下文趨勢意味著更多 HBM 需求(儘管 Flash 降低了單次計算的 HBM 訪問量,但更長序列帶來的總需求仍在增長) |
15 分鐘專家深入
核心思想:IO-Awareness
Tri Dao 等人在 2022 年發表的 FlashAttention 論文(“FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”)的核心洞察是:在現代 GPU 上,注意力的瓶頸不是 FLOPs,而是記憶體頻寬。
A100 的 Tensor Core 峰值 FP16 算力約 312 TFLOPS,但 HBM 頻寬僅約 2 TB/s。標準注意力需要將 N×N 的注意力矩陣寫入 HBM 再讀出來做 Softmax 和後續乘法,造成嚴重的頻寬瓶頸。
FlashAttention 的策略:永遠不把完整的 N×N 矩陣物化到 HBM 中。它將 Q、K、V 分成小塊(block),每個塊在 SRAM(片上快取記憶體,A100 每 SM 的共享記憶體最大可配置為 164 KB)中完成計算,利用線上 Softmax 技巧逐塊累加結果。
線上 Softmax(Online Softmax)
這是實現分塊計算的關鍵數學技巧(基於 Milakov & Gimelshein 2018 的工作):
- 標準 Softmax 需要知道整行的最大值和求和;
- 線上 Softmax 維護一個 執行中的最大值 m 和求和值 l,每處理一個新塊時:
- 如果發現新的最大值,對之前已累加的結果做 修正(乘以修正因子);
- 將當前塊的貢獻納入累加。
這保證了最終結果與標準注意力 數學等價(僅浮點舍入順序不同),不是近似。
FlashAttention-2 的三大改進
論文:“FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning”(Tri Dao, 2023)
改進一:減少非矩陣乘法 FLOPs
FA-1 在每個塊處理後需要做 rescaling(將之前累積的結果乘以修正因子),這引入了額外的逐元素操作。FA-2 將 rescaling 推遲到所有塊處理完之後統一進行,減少了迴圈內部的非 Tensor Core 操作。
改進二:跨序列維度並行
FA-1 的並行策略是:外層迴圈遍歷 K/V 塊(沿序列維度),內層迴圈遍歷 Q 塊。這意味著 Q 塊的處理是序列的。FA-2 反轉了迴圈順序——外層並行遍歷 Q 塊,內層序列遍歷 K/V 塊。這使得:
- 對於長序列,可以充分利用更多 SM;
- 序列並行度提高,不再受限於 batch_size × num_heads 的並行度。
改進三:更好的 Warp 間工作分配
A100 上每個 SM 有 4 個 warp 排程器(每個 SM 最多可駐留 64 個 warp,共 2048 個執行緒),每 warp 32 執行緒。FA-2 將 Q 塊在 warp 間按行分割,每個 warp 獨立處理部分行,減少 warp 間的共享記憶體通訊(synchronization 和 shared memory reads)。
結果(論文報告):
- 在 A100 SXM 上,FA-2 實現了約 230 TFLOPS(FP16/BF16),相當於 A100 Tensor Core 峰值(~312 TFLOPS)的約 72% 利用率 [FlashAttention-2 論文, Dao 2023];
- 相比 FA-1 提速約 2×;
- 相比標準 PyTorch attention 提速約 5–9×,記憶體節省約 5–20×(取決於序列長度)。
FlashAttention-3:擁抱 Hopper
論文:“FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision”(Shah, Dao et al., 2024)
FlashAttention-3 專門為 NVIDIA Hopper 架構(H100/H800)設計,利用了 Hopper 獨有的硬體特性:
特性一:非同步執行(Producer-Consumer 模式)
Hopper 引入了 Warp Group 概念(4 個 warp 組成一個 warpgroup,可以協作執行 WGMMA — Warp Group Matrix Multiply-Accumulate 指令)。FA-3 將 warpgroup 分為兩組:
┌─────────────────────────────────────────────┐
│ Warpgroup 0 (Producer) │
│ ┌──────────┐ ┌──────────┐ ┌──────────┐ │
│ │ TMA Load │ │ TMA Load │ │ TMA Load │ │ ← 非同步載入下一個塊
│ └──────────┘ └──────────┘ └──────────┘ │
│ ↓ 釋出 WGMMA 指令 │
├─────────────────────────────────────────────┤
│ Warpgroup 1 (Consumer) │
│ ┌──────────┐ │
│ │ Softmax │ ← 當 GEMM 在執行時,另一個 │
│ │ + Scaling│ warpgroup 做 Softmax │
│ └──────────┘ │
└─────────────────────────────────────────────┘
- TMA(Tensor Memory Accelerator):Hopper 專用硬體單元,可在計算進行時非同步搬運資料,不佔用 Tensor Core;
- WGMMA:新的 Tensor Core 指令,允許多個 warp 協作執行矩陣乘法,延遲更高但吞吐更大,且不阻塞呼叫 warp 的執行;
- 效果:Softmax 與 GEMM 重疊執行,幾乎消除 softmax 帶來的”氣泡”。
特性二:FP8 低精度 + Incoherent Processing
FP8(E4M3/E5M2)將 Tensor Core 吞吐翻倍,但注意力中的 Q、K 矩陣往往存在 outlier(異常值),導致 FP8 量化精度損失嚴重。
FA-3 引入 Incoherent Processing:在量化前對 Q、K 施加一個隨機正交變換(random orthogonal rotation),使得資料分佈更均勻、outlier 被”分散”,從而顯著改善 FP8 注意力的精度。這個變換是線上計算的,額外開銷很低。
特性三:塊級線上量化/反量化
FP8 的量化 scale 不是全域性固定的,而是 逐塊(per-block) 計算的。量化和反量化操作與計算流水線交錯執行,利用 TMA 和 WGMMA 的非同步特性掩蓋開銷。
結果(論文報告,H100 SXM):
- FP16:約 740 TFLOPS(論文報告值);
- FP8:約 1.2 PFLOPS(論文報告值);
- 相比 FA-2(在 H100 上執行)提速約 1.5–2×;
- FP16 模式下接近 H100 Tensor Core 的理論峰值利用率 [FA-3 論文, Shah et al. 2024]。
⚠️ 注意:上述 TFLOPS 數字均來自論文作者在特定配置(特定序列長度、batch size、head 數)下的測量,實際生產環境中的利用率會因模型配置不同而有所變化。
技術原理
標準注意力 vs FlashAttention 的計算與 IO 複雜度
標準注意力:
Input: Q ∈ R^{N×d}, K ∈ R^{N×d}, V ∈ R^{N×d}
Step 1: S = Q·K^T → 產出 N×N 矩陣,寫入 HBM
Step 2: P = softmax(S) → 讀 N×N、寫 N×N(HBM)
Step 3: O = P·V → 讀 N×N 和 V,寫 O(HBM)
HBM 訪問量: O(N²·d + N²) ≈ O(N²) (N² 項主導)
記憶體佔用: O(N²) (儲存完整注意力矩陣)
FLOPs: O(N²d) (兩個矩陣乘法)
FlashAttention 分塊計算:
將 Q 分成 T_r 個塊 (每塊大小 B_r × d)
將 K,V 分成 T_c 個塊 (每塊大小 B_c × d)
for i = 1 to T_r: ← 外層遍歷 Q 塊 (FA-2 的順序)
載入 Q_i 到 SRAM
初始化 O_i = 0, l_i = 0, m_i = -inf
for j = 1 to T_c: ← 內層遍歷 K/V 塊
載入 K_j, V_j 到 SRAM
計算 S_ij = Q_i · K_j^T (在 SRAM 中)
計算區域性 softmax:
m_ij = rowmax(S_ij)
P_ij = exp(S_ij - m_ij)
l_ij = rowsum(P_ij)
修正之前的結果:
m_new = max(m_i, m_ij)
l_new = l_i·exp(m_i - m_new) + l_ij·exp(m_ij - m_new)
O_i = O_i · (l_i·exp(m_i - m_new) / l_new)
+ (exp(m_ij - m_new) / l_new) · P_ij · V_j
m_i = m_new, l_i = l_new
O_i = O_i / l_i ← 最終歸一化
寫回 O_i 到 HBM
IO 複雜度對比 (M = SRAM 容量):
標準注意力: O(N²) HBM 訪問
FlashAttention: O(N²d² / M) HBM 訪問
當 d=128, M≈100KB 時, FA 的 HBM 訪問量約為標準實現的
d²/M ≈ 16384/100000 ≈ 1/6 (數量級估算,具體取決於分塊策略)
關鍵數值精度說明
FlashAttention 的核心保證是 數值等價性(numerical equivalence):
- 最終結果與標準 attention 的差距 僅來自浮點運算的結合順序差異;
- 不同於 LSH attention、稀疏 attention 等近似方法,FA 是 精確計算;
- FlashAttention-3 的 FP8 模式則引入了量化誤差,屬於近似計算,但通過 Incoherent Processing 將誤差控制在可接受範圍。
與硬體的對應關係
GPU 記憶體層次:
┌──────────────────────────────┐
│ HBM (視訊記憶體) │ A100: 80GB, 2 TB/s
│ Q, K, V, O 的主儲存位置 │
├──────────────────────────────┤
│ L2 Cache │ A100: 40 MB
├──────────────────────────────┤
│ Shared Memory (SRAM) │ A100: 每 SM 最高 164 KB
│ ← FlashAttention 的分塊 │ (可配置, 與 L1 共享)
│ 計算發生在這裡 │
├──────────────────────────────┤
│ Register File │ A100: 每 SM 256 KB
│ ← Tensor Core 直接操作 │
└──────────────────────────────┘
FlashAttention 的核心:
把計算"拉進" SRAM, 最小化 HBM 讀寫
技術演進史
| 時間 | 事件 | 意義 |
|---|---|---|
| 2018 | Milakov & Gimelshein 發表 Online Softmax | 分塊 Softmax 的數學基礎 |
| 2022.05 | Tri Dao 等發表 FlashAttention v1 | 首次將 IO-awareness 引入注意力計算,A100 上提速約 2× |
| 2022.10 | PyTorch 2.0 預覽版整合 FlashAttention | 走向標準化,使用者無需手動呼叫 CUDA 核心 |
| 2023.07 | FlashAttention-2 論文釋出 | 跨序列並行 + warp 最佳化,再提速約 2×,利用率 ~72% |
| 2023.Q4 | vLLM、TGI 等推論引擎預設啟用 | FlashAttention 成為推論標配 |
| 2024.07 | FlashAttention-3 論文釋出 | 適配 Hopper,引入非同步執行 + FP8,H100 上逼近峰值 |
| 2024.Q3+ | 各硬體平台跟進實現 | AMD Composable Kernel、Intel XPU 等均有類似實現 |
演進的核心驅動力
v1: "能不能不存 N×N 矩陣?" → tiling + online softmax
v2: "能不能讓 GPU 跑得更滿?" → 更好的並行與 warp 排程
v3: "能不能榨乾 Hopper 的新特性?" → 非同步執行 + 低精度
技術路線對比
FlashAttention 各版本橫向對比
| 維度 | FA-1 (2022) | FA-2 (2023) | FA-3 (2024) |
|---|---|---|---|
| 目標硬體 | A100/Ampere | A100/Ampere, 通用最佳化 | H100/Hopper 專用 |
| 數學等價性 | 精確 | 精確 | FP16 精確; FP8 近似(有精度保障) |
| A100 FP16 吞吐 | ~124 TFLOPS (論文報告) | ~230 TFLOPS (論文報告) | N/A (不針對 A100) |
| H100 FP16 吞吐 | N/A | ~230 TFLOPS (無 Hopper 最佳化) | ~740 TFLOPS (論文報告) |
| H100 FP8 吞吐 | N/A | N/A | ~1.2 PFLOPS (論文報告) |
| A100 Tensor Core 利用率 | ~40% [論文報告] | ~72% [論文報告] | N/A |
| 核心最佳化手段 | Tiling + Online Softmax | 序列並行 + Warp 最佳化 + 減少非 GEMM ops | TMA + WGMMA 非同步 + FP8 + Incoherent Proc. |
| 支援 head_dim | ≤ 256 | ≤ 256 | ≤ 256 (論文報告) |
| 支援 causal mask | 是 | 是 | 是 |
| 支援 GQA/MQA | 部分(後續版本) | 是 | 是 |
| 記憶體佔用 | O(N) | O(N) | O(N) |
與其他注意力加速方案對比
| 方案 | 型別 | 精度 | FLOPs | 典型加速 | 適用場景 |
|---|---|---|---|---|---|
| FlashAttention-2/3 | IO最佳化 | 精確(FA-3 FP8近似) | O(N²d) 不變 | 5–9× vs 標準 | 通用,首選 |
| 稀疏 Attention (Longformer等) | 稀疏化 | 近似 | O(N·w) | 取決於稀疏模式 | 超長序列 |
| Linear Attention (Performer等) | 核近似 | 近似 | O(N·d²) | 推論時友好 | 對精度容忍度高的場景 |
| Ring Attention | 分散式 | 精確 | O(N²d) | 多機長序列 | 多 GPU 超長訓練 |
| Sliding Window (Mistral) | 混合 | 近似 | O(N·w) | 固定視窗 | 推論/部分訓練 |
上下游
上游:硬體與基礎設施
┌─────────────────────────────────────────────────┐
│ 上游 │
├────────────────┬────────────────────────────────┤
│ GPU 硬體 │ NVIDIA A100/H100/H200/B100 │
│ │ AMD MI250X/MI300X (有對應實現) │
│ 記憶體層次 │ HBM 容量決定可處理的序列上限 │
│ CUDA 生態 │ Triton (編譯器), cuBLAS, CUTLASS │
│ Tensor Core │ FP16/BF16/FP8 矩陣乘法硬體 │
│ Hopper 專屬 │ TMA, WGMMA, Distributed Shared │
│ │ Memory, Thread Block Clusters │
└────────────────┴────────────────────────────────┘
中游:FlashAttention 本身
┌─────────────────────────────────────────────────┐
│ FlashAttention 核心 │
├────────────────┬────────────────────────────────┤
│ 核心演算法 │ Tiling + Online Softmax │
│ 實現語言 │ CUDA C++ + PTX 彙編(FA-3) │
│ 編譯器 │ Triton (Python DSL 實現的替代版) │
│ 配套庫 │ flash-attn (pip 安裝) │
│ 作者/維護 │ Tri Dao (Together AI / Princeton)│
└────────────────┴────────────────────────────────┘
下游:架構與應用
┌─────────────────────────────────────────────────┐
│ 下游 │
├────────────────┬────────────────────────────────┤
│ 訓練架構 │ PyTorch (SDPA flash backend) │
│ │ DeepSpeed, Megatron-LM, FSDP │
│ 推論引擎 │ vLLM, TensorRT-LLM, TGI, SGLang│
│ 模型庫 │ Hugging Face Transformers │
│ 模型 │ GPT-4, LLaMA 2/3, Mistral, │
│ │ Qwen, DeepSeek 等幾乎所有主流LLM│
│ 長上下文應用 │ 128K+ 文件理解, 程式碼庫分析 │
└────────────────┴────────────────────────────────┘
關鍵指標
| 指標 | 含義 | FlashAttention-2 典型值 | FlashAttention-3 典型值 |
|---|---|---|---|
| Tensor Core 利用率 | 實際 TFLOPS / 理論峰值 | ~72% (A100) [論文] | ~75%+ (H100 FP16) [論文] |
| HBM 訪問量 | 相比標準 attention 的減少比例 | 約 5–20× 減少 [取決於 seq_len] | 進一步減少(TMA 非同步) |
| 峰值吞吐 (FP16) | 單卡 attention kernel 吞吐 | ~230 TFLOPS (A100) [論文] | ~740 TFLOPS (H100) [論文] |
| 峰值吞吐 (FP8) | 單卡 attention kernel 吞吐 | N/A | ~1.2 PFLOPS (H100) [論文] |
| 記憶體複雜度 | attention 矩陣儲存 | O(N) | O(N) |
| FLOPs 複雜度 | 總計算量 | O(N²d)(不減少) | O(N²d)(不減少) |
| 支援序列長度 | 受限於視訊記憶體和 int32 索引 | 64K+ (A100 80GB) | 128K+ (H100 80GB) [估算] |
上述 TFLOPS 均為論文作者在最佳化配置下的測量值,生產環境中實際利用率因模型配置而異。
供需與市場資料
需求端驅動力
| 驅動力 | 資料點 |
|---|---|
| 上下文視窗增長 | 2K (GPT-3) → 128K (GPT-4 Turbo) → 1M+ (Gemini 1.5) |
| 模型規模增長 | 175B → 405B+ (LLaMA 3.1) → MoE 萬億引數 |
| 推論併發需求 | 生產部署中 batch size 持續增大 |
| 訓練成本 | 大型模型單次訓練成本 $10M–$100M+,2× 吞吐提升 = 數千萬美元節省 |
供給側格局
- NVIDIA 生態:FlashAttention 是事實標準,幾乎所有訓練都在 NVIDIA GPU 上完成
- AMD 跟進:ROCm 生態有
composable_kernel中的類 FlashAttention 實現,但成熟度和效能有差距 - 自研晶片:Google TPU 使用的是 JAX/XLA 內部最佳化的注意力實現(pallas flash attention);各 AI 晶片公司(Groq、Cerebras 等)需實現類似最佳化
- FPGA/ASIC:注意力 IO 最佳化的原理是通用的,但具體實現高度硬體相關
量化估算
- 全球 AI 訓練市場 2024 年約 $30B+(含 GPU 硬體 + 運營),FlashAttention 作為底層最佳化,其價值體現在 同等硬體下約 1.5–2× 的訓練效率提升(與標準實現對比)。
- 換言之,若沒有 FlashAttention 系列,行業可能需要 多投入 30–50% 的算力 才能達到當前的訓練吞吐 [粗略估算]。
代表公司與資本對映
| 公司/實體 | 與 FlashAttention 的關係 | 資本對映 |
|---|---|---|
| Together AI | Tri Dao 聯合創辦,FlashAttention 核心作者 | 私有輪估值約 $1.3B (2024 年報道) |
| NVIDIA (NVDA) | 核心硬體提供方,FA 極大提升了 GPU 利用率和賣卡邏輯 | NVDA |
| Princeton University | Tri Dao 所在機構,基礎研究 | N/A |
| Meta (META) | LLaMA 系列訓練重度依賴 FA,同時贊助相關研究 | META |
| Hugging Face | Transformers 庫深度整合 FA | 私有 |
| SK 海力士 (000660.KS) | HBM 主要供應商,FA 使能長上下文 → HBM 需求增長 | 000660.KS |
| Samsung (005930.KS) | HBM 供應商 | 005930.KS |
| AMD (AMD) | ROCm 生態跟進實現類似最佳化 | AMD |
Triton 編譯器路徑
值得注意的是,Tri Dao 團隊也在 OpenAI Triton(Python DSL GPU 編譯器)中實現了 FlashAttention。這意味著:
- 未來新硬體只需適配 Triton 後端,即可自動獲得 FlashAttention 級別的最佳化;
- 這對 NVIDIA 的 CUDA 護城河構成潛在的長期挑戰(降低硬體遷移成本)。
投資邏輯
核心投資邏輯
-
“AI 的 Amdahl 定律”:隨著模型規模和序列長度增長,注意力計算佔總計算的比例持續上升(尤其在長上下文場景),對注意力效率的最佳化將越來越關鍵。
-
標準鎖定效應:FlashAttention 已深度嵌入 PyTorch、Transformers、vLLM 等核心基礎設施,形成了事實標準,替代成本極高。
-
硬體協同演進:FA-3 表明,演算法最佳化與硬體架構是 共生關係——NVIDIA 每一代新架構都會催生新的演算法最佳化,而這些最佳化又進一步提升硬體的銷售價值。
-
長上下文是大趨勢:從 2K → 128K → 1M 的上下文視窗擴充套件趨勢不可逆轉,FlashAttention 是實現這一趨勢的 底層使能技術。
風險因素
| 風險 | 說明 |
|---|---|
| 新架構替代 | 線性 Attention、狀態空間模型(Mamba)等不使用標準注意力的架構可能減少對 FA 的依賴 |
| 硬體內建最佳化 | NVIDIA 可能在未來硬體中內建類似功能(如 Tensor Core 級別的注意力原語),降低 FA 的獨特價值 |
| 精度風險 | FP8 等低精度模式在某些任務上的精度損失尚未完全驗證 |
| 國產替代 | 中國 AI 晶片生態可能發展獨立的注意力最佳化路徑 |
常見誤讀糾偏
誤讀一:“FlashAttention 減少了注意力的 FLOPs”
糾偏:FlashAttention 不減少 FLOPs。注意力的計算量始終是 O(N²d),FA 一個 FLOP 都沒省。它減少的是 HBM 訪問量,瓶頸從頻寬(memory-bound)轉移到了計算(compute-bound),從而讓 Tensor Core 能真正被”餵飽”。
誤讀二:“FlashAttention 是一種近似注意力演算法”
糾偏:FlashAttention-1 和 FA-2 計算的是 精確注意力(exact attention),結果與標準實現數學等價(僅有浮點舍入順序差異)。不要把它和 Linformer、Performer、稀疏注意力等近似方法混為一談。FA-3 的 FP8 模式確實引入了量化近似,但 FP16/BF16 模式仍然是精確的。
誤讀三:“FlashAttention 讓模型能處理更長的上下文是因為它降低了記憶體複雜度”
部分糾偏:FA 將注意力矩陣的記憶體從 O(N²) 降到 O(N),這確實幫助了長序列。但注意——Q、K、V 矩陣本身仍需 O(N·d) 儲存,在超長序列(如 1M tokens)下,這部分的視訊記憶體消耗仍然是主要限制。真正的百萬級上下文還需要結合 序列並行(Ring Attention)、KV Cache 量化、分頁注意力(PagedAttention) 等技術。
誤讀四:“FA-3 只是 FA-2 加了個 FP8 支援”
糾偏:FP8 只是 FA-3 的特性之一。更關鍵的是 非同步執行架構——TMA 非同步搬運 + WGMMA 非同步計算 + warpgroup 間的 producer-consumer 模式,這使得 softmax 和 GEMM