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