注意力计算的数学很简单(QKᵀ → softmax → PV),但标准实现的瓶颈完全不在算力,而在显存 IO:中间矩阵 S(N×N)和 P 每次都要写回 HBM 再读出来。序列越长,这部分 IO 越致命。FlashAttention 用「分块计算 + 不落盘重算 + online softmax」把 HBM 访问从 O(N²) 降到 O(N),让长上下文推理第一次跑得动。本文讲透它的原理、变体、性能数据与手写内核要点。
前置:/ai-attention-optimization/(注意力计算与优化)、/ai-kernel-fusion-optimization/(算子融合)、/ai-cuda-memory-optimization/(显存管理)、/ai-vllm-system/(PagedAttention 与推理引擎)。
目录
- 1. 标准注意力的内存瓶颈:O(N^2) 的读写
- 2. FlashAttention 的思想:分块、重算与 IO 感知
- 3. Online Softmax:数值稳定的流式归一化
- 4. 前向与反向:重算如何省显存
- 5. 变体:Fused Attention、PagedAttention 与 MLA
- 6. 算子融合:FlashAttention 与其余算子
- 7. 性能数据:H100 与 A100 上的实测
- 8. 手写内核要点:CUDA 分块与寄存器
- 9. 选型与踩坑
- 10. 速查表与一句话记忆
- 延伸阅读
1. 标准注意力的内存瓶颈:O(N^2) 的读写
标准注意力三步:
Q, K, V ∈ R^(N×d)
S = Q @ K^T # 形状 N×N
P = softmax(S) # 还是 N×N
O = P @ V # 形状 N×d
内存特征:
□ S 和 P 都是 N×N 的中间矩阵 → 必须写回 HBM
□ 大矩阵乘法要切块循环 → 中间结果反复写读 HBM
□ N 越大,S/P 越大,HBM 带宽成为真正的瓶颈
为什么「算力够用但跑不快」:
□ A100 HBM 带宽 ~2 TB/s,算力 312 TFLOPS
□ 注意力是「带宽受限」算子:IO 决定墙钟时间
□ 即使算力空闲,数据搬不动 → 长序列卡在 IO
IO 量估算:
S 写回 HBM + 读回:2 × N² × 4 bytes(FP32)
P 写回 + 读回:2 × N² × 4 bytes
N=4096:中间矩阵 IO ≈ 256 MB(一次注意力)
N=32768:≈ 16 GB → 光搬数据就要 8 秒(A100)
→ 这就是「长上下文跑不动」的直接原因
工程要点:标准注意力的问题不是 FLOPs,而是把 N×N 的中间矩阵 S、P 反复写读 HBM——序列越长 IO 越致命。算力再强也被带宽卡死,所以优化方向从「减少计算」转向「减少 IO」。
2. FlashAttention 的思想:分块、重算与 IO 感知
FlashAttention 的核心命题:让中间矩阵不落 HBM。
三个关键手段:
□ 分块(Tiling):
- 把 Q/K/V 切成小块,一次只算一块注意力
- 小块计算结果留在 SRAM(片上),不写回 HBM
□ 重算(Recompute):
- 反向传播不保存 S、P(省显存)
- 反向时用保存的 Q/K/V 重新算一遍
□ IO 感知(IO-aware):
- 目标从「减少 FLOPs」变成「减少 HBM 访问」
- 块大小按 SRAM 容量决定,最大化片上复用
HBM 访问对比:
□ 标准:O(N²) 次 HBM 访问
□ FlashAttention:O(N²×d²/SRAM) 次 → 长序列下近似 O(N)
分块示意图:
Q 切 [Q₁, Q₂, ...],K/V 也切块
对每个 Q 块:
加载 Qᵢ、遍历所有 K/V 块
在 SRAM 里算局部注意力、累积 O 与 softmax 统计量
块间结果在片上合并,直到最后才写一次 O
工程要点:FlashAttention 用「分块让中间结果留在 SRAM + 反向重算省显存 + 以 HBM 访问数为优化目标」三招,把注意力从「带宽受限」变回「算力受限」——这是长上下文能跑起来的根本原因。
3. Online Softmax:数值稳定的流式归一化
Softmax 需要「全局 max 与全局 sum」,但分块计算一次只见局部——online softmax 解决了流式归一化问题。
标准 softmax:
□ 需要完整一行算 max,再算 exp,再归一化
□ 要求整行 S 都在手边 → 与分块矛盾
Online Softmax 技巧:
□ 每个块维护「局部 max(m)」与「局部 sum(l)」
□ 新块到来时更新:
- 新 max = max(旧 max, 新块 max)
- 用新 max 修正旧块的累积值(rescaling)
□ 保证数值稳定(每步减 max)同时不落盘
修正公式(直觉):
旧块统计量按「新 max」重新缩放
l_new = l_old × exp(m_old - m_new) + l_block
m_new = max(m_old, m_block)
流式过程:
块1:max=2.0, sum=3.1
块2:max=3.5 → 修正块1:sum_old × exp(2.0-3.5)
总 max=3.5,总 sum = 修正后 + 块2
→ 最终 softmax 分母 = 总 sum,分子逐块用全局 max 归一
工程要点:Online Softmax 是分块注意力的数学基石——每个块只维护局部 max/sum,新块到来时用 rescaling 修正历史统计,全程数值稳定且不需要完整矩阵在内存里。没有它,分块和归一化无法共存。
4. 前向与反向:重算如何省显存
FlashAttention 的显存收益主要来自反向传播不保存 S、P。
前向:
□ 只保存输出 O 与每块的 softmax 统计量(m、l)
□ 不保存 S、P → 峰值显存从 O(N²) 降到 O(N)
□ 块内计算在 SRAM,最后写 O
反向:
□ 标准实现需要 S、P 算梯度 → 全存显存
□ FlashAttention 反向「重算」:
- 用保存的 Q/K/V + 统计量,重新算一遍注意力
- 重算的 FLOPs 多一点点,但省下 O(N²) 显存
□ 重算在训练/长上下文微调中同样关键
显存对比:
□ 标准:注意力中间矩阵 O(N²)
□ FlashAttention:O(N)(只存 Q/K/V/O + 统计量)
反向重算流程:
加载 Q/K/V 块 → 重新算 S、P(不落盘,在 SRAM)
→ 结合上游梯度 dO 算 dQ/dK/dV
→ 每个块算完即弃,峰值显存只与「单块」有关
工程要点:FlashAttention 用「前向只存统计量 + 反向重算注意力」把注意力显存从 O(N²) 压到 O(N)——用一点额外 FLOPs 换巨额显存。这正是超长上下文训练/推理在显存上「跑得动」的关键。
5. 变体:Fused Attention、PagedAttention 与 MLA
FlashAttention 衍生出一系列变体,各自解决不同场景的问题。
FlashAttention 变体族谱:
□ FlashAttention-2:进一步优化并行策略与分块
- 在 H100/A100 上达到接近峰值算力利用率
□ FlashDecoding:解码阶段的变体
- 针对「一个 query 对一个长 KV」的 decode 特征
- 把 KV 切块并行算部分和,再 reduce
□ PagedAttention(vLLM):
- FlashAttention 的显存管理 + 页式 KV 存储
- 注意力计算支持「非连续块」的 KV 读取
□ 分页 FlashAttention(Paged + Flash 融合):
- 兼顾页式显存复用与 IO 优化
□ MLA(Multi-head Latent Attention):
- DeepSeek 的低秩 KV 压缩
- 把 KV 压缩成潜在向量 → KV Cache 大幅缩小
- 推理时再解压(需要内核级支持)
变体选择场景:
短上下文高吞吐 → 标准 FlashAttention-2
超长上下文 → FlashAttention + 页式/稀疏
解码带宽受限 → FlashDecoding(KV 并行分块)
显存极度紧张 → MLA 低秩压缩
工程要点:FlashAttention 已从「一个内核」长成一个「注意力内核家族」——FlashAttention-2 拉满算力、FlashDecoding 优化 decode、PagedAttention 加页式显存管理、MLA 用低秩压缩 KV。选型取决于你的瓶颈是算力、IO、显存还是 KV 容量。
6. 算子融合:FlashAttention 与其余算子
FlashAttention 与融合优化不是孤立的,它经常和相邻算子一起被融合。
融合的价值:
□ 每少一次 HBM 写读,就少一段 IO 时间
□ 注意力的上下游算子可以「粘」进同一个 kernel
常见融合组合:
□ Attention + 残差 + LayerNorm:
- pre-norm:Q/K/V 投影 → 注意力 → 残差 → LayerNorm
- 合成一个 kernel,中间结果不落盘
□ QKV 投影融合:把三个线性层合成一个矩阵乘
□ 输出投影 + 残差 + MLP 的连续融合
□ RoPE 位置编码融合进 Q/K 投影
与 MoE / MoA 的配合:
□ 注意力部分用 FlashAttention(IO 优化)
□ 专家部分单独内核(序列专家 GEMM)
→ 各算各的最优,避免强行融合反而退化
融合前后对比:
未融合:Q投影→写HBM → 注意力→写HBM → 残差→写HBM
融合后:Q投影+注意力+残差+LN 一个 kernel 完成
→ HBM 写读次数从 4 次降到 1 次
工程要点:FlashAttention 是算子融合的「中心节点」——把 QKV 投影、RoPE、注意力、残差、LayerNorm 融成一个 kernel,HBM 访问次数成倍下降。但融合要「该融才融」:MoE 专家计算与注意力保持分离,各用最优内核。
7. 性能数据:H100 与 A100 上的实测
注意力内核优化的目标,是把「带宽受限」推回「算力受限」。
衡量指标:
□ HBM 访问量(越低越好)
□ 算力利用率(越接近峰值越好)
□ 端到端注意力耗时
FlashAttention 实测特征(社区公开数据口径):
□ A100 上 FlashAttention 相比标准实现提速 2-4x
□ 长序列收益更大:序列越长,IO 优势越明显
□ FlashAttention-2 在 A100/H100 达到 ~70%-90% 算力利用率
□ FlashDecoding 在 decode 阶段比标准快数倍
瓶颈回归点:
□ 短序列(N<256):算力与 IO 差距小,收益收窄
□ 极长序列:从「IO 受限」彻底变成「算力受限」
典型量级(相对值,供直觉参考):
N=512 :Flash ≈ 1.5-2x
N=2048 :Flash ≈ 2-3x
N=8192 :Flash ≈ 3-4x
N=32768 :Flash 收益进一步拉大(IO 差 O(N²) vs O(N))
工程要点:FlashAttention 的收益随序列长度放大——短序列只有 1.5-2x,长序列可到 3-4x 以上,H100/A100 上能逼近算力峰值。评估注意力内核的性能,核心看「HBM 访问量」与「算力利用率」两个数。
8. 手写内核要点:CUDA 分块与寄存器
如果你想写自己的注意力内核,关键要点如下。
手写内核的核心决策:
□ 分块大小:由 SRAM 容量决定(如 128×128)
- 太大:SRAM 溢出 → 被迫落盘
- 太小:块间开销高、复用差
□ 数据布局:Q/K/V 按块连续加载,支持张量核心
□ online softmax 统计量:每个线程块维护 m/l
□ 反向重算:不保存 S/P,用统计量重建
CUDA 层面:
□ 用共享内存(shared memory)模拟 SRAM 手动分块
□ 寄存器级切分:让每个线程处理一小块,最大化复用
□ 用 Tensor Core / cuBLAS 风格的矩阵乘原语
□ 注意 bank conflict:共享内存布局要对齐
伪代码骨架(前向,一块):
load Q_block, K_block, V_block → shared memory
S_block = Q_block @ K_block^T # 在片上
m_new = max(m_old, rowmax(S_block))
P_block = exp(S_block - m_new)
l_new = l_old * exp(m_old - m_new) + rowsum(P_block)
O_block = rescale(O_old) + P_block @ V_block
m_old, l_old, O_old = m_new, l_new, O_block
工程要点:手写注意力内核的要点是「SRAM 分块 + online softmax + 反向重算」三件套——块大小由 SRAM 定,统计量放线程块内,张量核心做矩阵乘。工程上先跑通数值(与标准实现对比),再谈优化到算力上限。
9. 选型与踩坑
注意力内核的选择与常见坑。
选型建议:
□ 生产推理引擎:直接用 vLLM/SGLang/TensorRT-LLM 内置内核
□ 需要特殊变体:FlashDecoding、Paged、MLA 看引擎支持
□ 自研内核:先跑通数学,再优化,再对比标准实现
□ 训练侧:PyTorch 的 flash_attn 库或 xformers 内存高效注意力
高频踩坑:
□ 直接手写而不对比数值 → 精度偏差(online softmax 细节错)
□ 分块大小拍脑袋 → 块太大落盘、太小复用差
□ 忽视 bank conflict → 共享内存命中率暴跌
□ 短序列强上 Flash → 收益小、调度开销反而拖慢
□ 把 FLOPs 当性能指标 → 真正瓶颈是 HBM 访问量
□ 反向重算的 FLOPs 上升没算进成本 → 训练侧意外变慢
数值验证:
□ 前向:输出与 float64 参考一致(相对误差 < 1e-5 量级)
□ 反向:梯度与标准实现一致(gradcheck)
验证清单:
□ max 值稳定性(数值溢出边界测试)
□ 掩码/因果注意力支持
□ 不同 dtype(FP16/BF16/FP8)的精度
□ 变长序列(分页/打包)支持
工程要点:注意力内核选型的默认动作是「用成熟引擎的内置实现」——只有特殊变体才值得自研。踩坑集中在数值细节(online softmax 实现错一处就不一致)、分块大小和「拿 FLOPs 当指标」。自研必须过数值对比和 gradcheck。
10. 速查表与一句话记忆
| 问题 | 一句话答案 |
|---|---|
| 注意力真正瓶颈 | 中间矩阵 S/P 反复写读 HBM,不是算力 |
| FlashAttention 三招 | 分块(SRAM 片上)+ 重算(反向省显存)+ IO 感知 |
| HBM 访问降多少 | 从 O(N²) 降到 O(N),序列越长收益越大 |
| Online Softmax 解决什么 | 分块下的流式归一化(局部 max/sum + rescaling) |
| 反向怎么省显存 | 不存 S/P,用保存的 Q/K/V + 统计量重算 |
| 有哪些变体 | FlashAttention-2、FlashDecoding、Paged、MLA |
| 和谁融合 | QKV 投影、RoPE、残差、LayerNorm 一个 kernel |
| 提速多少 | 短序列 1.5-2x,长序列 3-4x 以上 |
| 自研核心决策 | 分块大小按 SRAM、online softmax、反向重算 |
| 第一性能指标 | HBM 访问量(不是 FLOPs) |
一句话记忆:标准注意力被 HBM 带宽卡死(S/P 反复写读 O(N²));FlashAttention = 分块留 SRAM + online softmax 流式归一 + 反向重算省显存(HBM 降 O(N));变体按瓶颈选(Flash-2 拉算力/FlashDecoding 优化 decode/MLA 压 KV)——HBM 访问量才是第一指标。
延伸阅读
- /ai-attention-optimization/ — 注意力计算与 KV Cache 优化
- /ai-kernel-fusion-optimization/ — 算子融合与推理内核
- /ai-cuda-memory-optimization/ — GPU 显存管理与分块
- /ai-vllm-system/ — PagedAttention 与推理引擎实现
- /ai-inference-engine-comparison/ — 各引擎内核实现对比
- 高性能计算专题 — GPU 内核与 CUDA 优化
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。