模型层 开放阅读

FlashAttention-2/3

FlashAttention-2/3

概念 ID
flashattention-2-3
更新时间
2026-05-29
来源数量
待补

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 提速约
  • 相比标准 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 读写

技术演进史

时间事件意义
2018Milakov & Gimelshein 发表 Online Softmax分块 Softmax 的数学基础
2022.05Tri Dao 等发表 FlashAttention v1首次将 IO-awareness 引入注意力计算,A100 上提速约 2×
2022.10PyTorch 2.0 预览版集成 FlashAttention走向标准化,用户无需手动调用 CUDA 内核
2023.07FlashAttention-2 论文发布跨序列并行 + warp 优化,再提速约 2×,利用率 ~72%
2023.Q4vLLM、TGI 等推理引擎默认启用FlashAttention 成为推理标配
2024.07FlashAttention-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/AmpereA100/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/AN/A~1.2 PFLOPS (论文报告)
A100 Tensor Core 利用率~40% [论文报告]~72% [论文报告]N/A
核心优化手段Tiling + Online Softmax序列并行 + Warp 优化 + 减少非 GEMM opsTMA + WGMMA 异步 + FP8 + Incoherent Proc.
支持 head_dim≤ 256≤ 256≤ 256 (论文报告)
支持 causal mask
支持 GQA/MQA部分(后续版本)
内存占用O(N)O(N)O(N)

与其他注意力加速方案对比

方案类型精度FLOPs典型加速适用场景
FlashAttention-2/3IO优化精确(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 AITri Dao 联合创办,FlashAttention 核心作者私有轮估值约 $1.3B (2024 年报道)
NVIDIA (NVDA)核心硬件提供方,FA 极大提升了 GPU 利用率和卖卡逻辑NVDA
Princeton UniversityTri Dao 所在机构,基础研究N/A
Meta (META)LLaMA 系列训练重度依赖 FA,同时赞助相关研究META
Hugging FaceTransformers 库深度集成 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 护城河构成潜在的长期挑战(降低硬件迁移成本)。

投资逻辑

核心投资逻辑

  1. “AI 的 Amdahl 定律”:随着模型规模和序列长度增长,注意力计算占总计算的比例持续上升(尤其在长上下文场景),对注意力效率的优化将越来越关键。

  2. 标准锁定效应:FlashAttention 已深度嵌入 PyTorch、Transformers、vLLM 等核心基础设施,形成了事实标准,替代成本极高。

  3. 硬件协同演进:FA-3 表明,算法优化与硬件架构是 共生关系——NVIDIA 每一代新架构都会催生新的算法优化,而这些优化又进一步提升硬件的销售价值。

  4. 长上下文是大趋势:从 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

source: 公开披露与公开资料整理 本页仅用于产业链学习、信息检索和研究辅助;不构成投资建议,不预测涨跌,不提供买卖、仓位或目标价建议。
完整概念页 复盘 13 节结构 公司投研页 沿产业链找到受益公司 投资课 把概念转成可跟踪模型