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 数据流,伪代码只用于说明结构。
- 公式
- 流程
- 伪代码结构
FlashAttention 算法流程
┌───────────────────────────────────────────────────────────────┐
│ FlashAttention 单块计算流程 │
├───────────────────────────────────────────────────────────────┤
│ │
│ 输入:Q 块 (Bc×d), K 块序列 {Kj}, V 块序列 {Vj} │
│ 输出:O 块 (Bc×d) │
│ │
│ 1. 加载 Q 块到 SRAM │
│ 2. 初始化 m = -∞, ℓ = 0, O = 0 │
│ │
│ 3. For each K/V 块 j: │
│ ┌────────────────────────────────────────────┐ │
│ │ a. 加载 Kj, Vj 到 SRAM │ │
│ │ b. 计算 Sij = Qi × Kj^T (Bc×Bc 矩阵) │ │
│ │ c. 计算 mij = max(m, rowmax(Sij)) │ │
│ │ d. 更新 ℓ = exp(m-mij)×ℓ + rowsum(exp(Sij-mij)) │ │
│ │ e. 更新 O = diag(exp(m-mij))×O + exp(Sij-mij)×Vj │ │
│ │ f. 更新 m = mij │ │
│ └────────────────────────────────────────────┘ │
│ │
│ 4. 归一化 O = O / ℓ │
│ 5. 写回 O 块到全局内存 │
│ │
└───────────────────────────────────────────────────────────────┘
FlashAttention 代码结构
以下代码是用于说明数据流的伪代码骨架,省略了线程映射、向量化加载、边界处理、矩阵指令调用和逐行 softmax 状态等实现细节,不能作为独立 kernel 直接编译。
template<int Bc, int Br, int d>
__global__ void flashAttention(
float* Q, float* K, float* V, float* O,
int N, int d_model
) {
// 共享内存
__shared__ float Q_shared[Bc][d];
__shared__ float K_shared[Br][d];
__shared__ float V_shared[Br][d];
// 寄存器状态
float m_i = -INFINITY; // running max
float ell_i = 0.0f; // 运行分母(running denominator)
float O_accum[Bc][d] = {0.0f}; // 累加器
// Block 索引
int bx = blockIdx.x;
int by = blockIdx.y;
// 加载 Q 块到 SRAM
loadQBlock(Q, bx, by, Q_shared);
__syncthreads();
// 主循环:遍历所有 K/V 块
for (int j = 0; j < (N + Br - 1) / Br; j++) {
// 加载 K, V 块
loadKVBlock(K, V, j, K_shared, V_shared);
__syncthreads();
// 计算 Sij = Qi × Kj^T
float S[Bc][Br];
computeS(Q_shared, K_shared, S);
// online softmax 更新
float m_new = -INFINITY;
float P[Bc][Br];
// 计算新的最大值
for (int i = 0; i < Bc; i++) {
for (int jj = 0; jj < Br; jj++) {
m_new = max(m_new, S[i][jj]);
}
}
m_new = max(m_i, m_new);
// 计算概率并更新累加器
float scale = exp(m_i - m_new);
ell_i = ell_i * scale;
for (int i = 0; i < Bc; i++) {
for (int jj = 0; jj < Br; jj++) {
P[i][jj] = exp(S[i][jj] - m_new);
ell_i += P[i][jj];
}
}
// 更新 O = diag(scale) × O + P × V
updateO(O_accum, P, V_shared, scale);
// 更新状态
m_i = m_new;
__syncthreads();
}
// 归一化
for (int i = 0; i < Bc; i++) {
for (int j = 0; j < d; j++) {
O_accum[i][j] /= ell_i;
}
}
// 写回结果
storeOBlock(O, bx, by, O_accum);
}
FlashAttention 实现细节
实现时通常先确定块大小,再设计数据搬运流水线,最后用编译器报告和性能分析工具检查寄存器压力。下面三个标签页对应这三个调优入口。
- 块大小
- 软件流水线
- 寄存器压力
块大小(Block Size)选择
关键约束:共享内存容量限制
对于 MP31 架构(单 MP 192KB 共享内存):
典型配置(d=128, FP16):
| 参数 | 值 | 说明 |
|---|---|---|
| 256 | Q 块大小 | |
| 128 | K/V 块大小 | |
| 共享内存使用 | ~192KB | 接近上限 |
不同隐藏维度(Head Dimension)
| d | 推荐 Bc | 推荐 Br | Q/K/V 基础 SRAM 使用 |
|---|---|---|---|
| 64 | 512 | 256 | 约 128KB |
| 128 | 256 | 128 | 约 128KB |
| 256 | 128 | 64 | 约 128KB |
实际共享内存使用还取决于双缓冲、online softmax 状态和其他临时缓冲,需要结合具体 kernel 实现确认。
软件流水线设计
┌───────────────────────────────────────────────────────────────┐
│ FlashAttention 软件流水线(双缓冲) │
├───────────────────────────────────────────────────────────────┤
│ │
│ 迭代 0: [Load Q] → [Load K0,V0] → [Compute QK0] → [Update O0]│
│ │
│ 迭代 1: [Load K1,V1] → [Compute QK1] → [Update O1] │
│ ↑ │
│ └── 与上一次计算重叠 │
│ │
│ 迭代 2: [Load K2,V2] → [Compute QK2] │
│ ↑ │
│ └── 与上一次计算重叠 │
│ │
└───────────────────────────────────────────────────────────────┘
优化技巧:
- 双缓冲:使用两组共享内存,交替加载和计算
- 张量内存引擎(TME)
- 预计算 QK:第一轮预先计算,主循环内 overlap softmax
寄存器压力管理
主要寄存器消耗:
- 累加器: 个 FP32 寄存器
- 概率矩阵: 个 FP32 寄存器
- 临时状态: tile、online softmax 的 running max/denominator 等
优化策略:
- 减小 :降低寄存器压力(与 SRAM 约束冲突)
- 减小每个线程负责的输出元素数量:降低 累加器占用
- 检查编译器寄存器使用报告,避免寄存器 spill 到局部内存
FlashAttention-2/3 改进
- FlashAttention-2
- FlashAttention-3
显存与性能收益趋势
理论显存占用对比(FP16)
| 序列长度 | 标准 Attention | FlashAttention | 降低倍数 |
|---|---|---|---|
| N=512 | 4MB | 0.5MB | 8x |
| N=2048 | 68MB | 2MB | 34x |
| N=4096 | 268MB | 4MB | 67x |
| N=16384 | 4.3GB | 16MB | 268x |
性能收益趋势(FP16)
| 序列长度 | 标准 Attention | FlashAttention | 说明 |
|---|---|---|---|
| N=512 | 访存开销较低 | 收益有限 | 短序列下 kernel 开销占比更高 |
| N=2048 | 注意力矩阵访存压力增大 | 收益明显 | 分块和重计算开始体现优势 |
| N=4096 | O(N²) 显存与访存压力明显 | 收益显著 | 避免落地完整注意力矩阵 |
| N=16384 | 显存与访存压力很高 | 更适合长序列 | 实际性能需以目标硬件实测为准 |
上表用于说明随序列长度增长的收益趋势,不代表固定硬件或固定 kernel 的性能承诺。