FlashAttention 优化
FlashAttention 优化核心原则
- 避免存储完整的注意力矩阵(O(N²) 显存)
- online softmax(在线 Softmax)消除全局依赖
- 分块(Tiling)计算,充分利用片上 SRAM
- 重计算(Recompute)换存储
FlashAttention 是一种 IO-aware attention 算法。它不改变 self-attention 的数学结果,而是通过分块、online softmax 和重计算减少对 HBM 的读写。
| 读者目标 | 建议先看 |
|---|---|
| 理解为什么需要 FlashAttention | 标准 Self-Attention 的问题、FlashAttention 核心创新 |
| 理解算法如何工作 | online softmax、FlashAttention 完整算法 |
| 做 kernel 参数调优 | 块大小选择、软件流水线、寄存器压力管理 |
标准 Self-Attention 的问题
标准 self-attention 会显式生成 的注意力矩阵。序列越长,注意力矩阵带来的显存和访存压力越明显,这也是 FlashAttention 主要要解决的问题。
- 计算流程
- 显存压力
标准 Attention 公式
其中:
- 是 矩阵(N 为序列长度,d 为隐藏维度)
- 输出 也是 矩阵
标准实现的三步计算
Step 1: S = Q × K^T (N × N 注意力矩阵)
Step 2: P = softmax(S) (N × N 概率矩阵)
Step 3: O = P × V (N × d 输出矩阵)
显存需求分析
| 矩阵 | Shape | 显存占用(FP16) |
|---|---|---|
| Q, K, V | N × d | 3 × 2N × d Bytes |
| S (注意力矩阵) | N × N | 2N² Bytes |
| P (概率矩阵) | N × N | 2N² Bytes |
| O | N × d | 2N × d Bytes |
| 总计 | - | 4N² + 8Nd Bytes |
问题:当 N 较大时,N² 项主导显存占用
示例:N=4096, d=128
- Q, K, V: 3 × 2 × 4096 × 128 = 3MB
- S, P: 4 × 4096² = 128MB(占 97%!)
长序列的显存瓶颈
| 序列长度 | S+P 显存占用 | 是否可接受 |
|---|---|---|
| N=512 | 2MB | 可接受 |
| N=2048 | 32MB | 需关注 |
| N=4096 | 128MB | 显存压力明显增加 |
| N=16384 | 2GB | 长序列场景需重点评估 |
FlashAttention 核心创新
FlashAttention 概述
FlashAttention 的核心不是近似计算,而是改变 attention kernel 的数据流。它把 Q/K/V 分块加载到片上 SRAM,在块内完成矩阵乘、softmax 更新和输出累加,避免将完整的 和 矩阵写回全局内存。
| 创新点 | 作用 | 效果 |
|---|---|---|
| 分块(Tiling) | 将大矩阵分解为小块 | 控制片上 SRAM 占用 |
| Recompute(重计算) | 避免存储完整注意力矩阵 | 显存降低 O(N²) → O(N) |
| IO 感知设计 | 优化 GPU 内存层次数据流 | 带宽利用率提升 |
关键挑战:Softmax 的全局依赖
Softmax 公式:
问题:分母需要全局求和,无法直接分块计算
解决方案:online softmax(在线 Softmax)
online softmax(在线 Softmax)
标准 Softmax 回顾
# 标准 Softmax(需要三次循环)
# 计算最大值(用于数值稳定)
m = max(S)
# 计算分母(需要第一次循环结束后才能开始)
Z = sum(exp(S - m))
# 计算输出(需要第二次循环结束后才能开始)
P = exp(S - m) / Z
问题:三次循环,无法融合,访存效率低
online softmax(在线 Softmax)核心思想
关键洞察:推导 和 的递归公式,消除全局依赖
定义:
- (前 j 项最大值)
- (归一化分母)
递归关系:
online softmax(在线 Softmax)算法
# online softmax(单次循环)
m = -inf # 运行最大值(running maximum)
ell = 0.0 # 运行分母(running denominator)
for j in range(N):
# 更新最大值和分母
m_new = max(m, S[j])
ell_new = exp(m - m_new) * ell + exp(S[j] - m_new)
# 更新输出(递 归调整之前的值)
for i in range(j + 1):
P[i] = P[i] * exp(m - m_new)
P[j] = exp(S[j] - m_new) / ell_new
# 更新状态
m = m_new
ell = ell_new
优势:
- 单次循环即可完成
- 中间状态 可保存在 SRAM 中
- 无需存储完整的 S 矩阵
online softmax(在线 Softmax)数值稳定性
Safe Softmax 技巧:
其中 ,避免数值溢出。
online softmax 天然支持:
- 递归公式中已包含减最大值操作
- 无需额外处理
FlashAttention 完整算法
FlashAttention 的执行路径可以先按“加载 Q 块、遍历 K/V 块、在线更新 softmax 状态、写回输出”来理解。公式适合核对数学等价性,流程图适合理解 kernel 数据流,伪代码只用于说明结构。