Skip to main content

FlashAttention 优化

FlashAttention 优化核心原则
  1. 避免存储完整的注意力矩阵(O(N²) 显存)
  2. online softmax(在线 Softmax)消除全局依赖
  3. 分块(Tiling)计算,充分利用片上 SRAM
  4. 重计算(Recompute)换存储

FlashAttention 是一种 IO-aware attention 算法。它不改变 self-attention 的数学结果,而是通过分块、online softmax 和重计算减少对 HBM 的读写。

读者目标建议先看
理解为什么需要 FlashAttention标准 Self-Attention 的问题、FlashAttention 核心创新
理解算法如何工作online softmax、FlashAttention 完整算法
做 kernel 参数调优块大小选择、软件流水线、寄存器压力管理

标准 Self-Attention 的问题

标准 self-attention 会显式生成 N×NN \times N 的注意力矩阵。序列越长,注意力矩阵带来的显存和访存压力越明显,这也是 FlashAttention 主要要解决的问题。

标准 Attention 公式

Attention(Q,K,V)=softmax(QKTd)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d}}\right)V

其中:

  • Q,K,VQ, K, VN×dN \times d 矩阵(N 为序列长度,d 为隐藏维度)
  • 输出 OO 也是 N×dN \times d 矩阵

标准实现的三步计算

Step 1: S = Q × K^T (N × N 注意力矩阵)
Step 2: P = softmax(S) (N × N 概率矩阵)
Step 3: O = P × V (N × d 输出矩阵)

FlashAttention 核心创新

FlashAttention 概述

FlashAttention 的核心不是近似计算,而是改变 attention kernel 的数据流。它把 Q/K/V 分块加载到片上 SRAM,在块内完成矩阵乘、softmax 更新和输出累加,避免将完整的 SSPP 矩阵写回全局内存。

创新点作用效果
分块(Tiling)将大矩阵分解为小块控制片上 SRAM 占用
Recompute(重计算)避免存储完整注意力矩阵显存降低 O(N²) → O(N)
IO 感知设计优化 GPU 内存层次数据流带宽利用率提升

关键挑战:Softmax 的全局依赖

Softmax 公式

Pi=eSijeSjP_i = \frac{e^{S_i}}{\sum_j e^{S_j}}

问题:分母需要全局求和,无法直接分块计算

解决方案: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\ell_jmjm_j 的递归公式,消除全局依赖

定义

  • mj=max0ijSim_j = \max_{0 \leq i \leq j} S_i(前 j 项最大值)
  • j=i=0jeSimj\ell_j = \sum_{i=0}^j e^{S_i - m_j}(归一化分母)

递归关系

mj=max(mj1,Sj)m_j = \max(m_{j-1}, S_j) j=emj1mjj1+eSjmj\ell_j = e^{m_{j-1} - m_j} \cdot \ell_{j-1} + e^{S_j - m_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

优势

  • 单次循环即可完成
  • 中间状态 m,m, \ell 可保存在 SRAM 中
  • 无需存储完整的 S 矩阵

online softmax(在线 Softmax)数值稳定性

Safe Softmax 技巧

Pi=eSimjeSjmP_i = \frac{e^{S_i - m}}{\sum_j e^{S_j - m}}

其中 m=max(S)m = \max(S),避免数值溢出。

online softmax 天然支持

  • 递归公式中已包含减最大值操作
  • 无需额外处理

FlashAttention 完整算法

FlashAttention 的执行路径可以先按“加载 Q 块、遍历 K/V 块、在线更新 softmax 状态、写回输出”来理解。公式适合核对数学等价性,流程图适合理解 kernel 数据流,伪代码只用于说明结构。

FlashAttention 公式推导

结合 online softmax 后的 Attention

对于每个输出块 Oi:1. 加载 Qi (Bc×d)2. 初始化 mi=,i=0,Oi=03. 对于每个 K, V 块 j:a. 加载 Kj,Vjb. 计算 Sij=QiKjTc. 计算 mij=max(mi,rowmax(Sij))d. 更新 i=emimiji+rowsum(eSijmij)e. 更新 Oi=diag(emimij)Oi+eSijmijVjf. 更新 mi=mij4. 归一化 Oi=Oi/i\begin{aligned} &\text{对于每个输出块 } O_i: \\ &1.\ \text{加载 } Q_i \text{ (B}_c \times d\text{)} \\ &2.\ \text{初始化 } m_i = -\infty, \ell_i = 0, O_i = 0 \\ &3.\ \text{对于每个 K, V 块 } j: \\ &\quad \text{a. 加载 } K_j, V_j \\ &\quad \text{b. 计算 } S_{ij} = Q_i K_j^T \\ &\quad \text{c. 计算 } m_{ij} = \max(m_i, \text{rowmax}(S_{ij})) \\ &\quad \text{d. 更新 } \ell_i = e^{m_i - m_{ij}} \ell_i + \text{rowsum}(e^{S_{ij} - m_{ij}}) \\ &\quad \text{e. 更新 } O_i = \text{diag}(e^{m_i - m_{ij}}) O_i + e^{S_{ij} - m_{ij}} V_j \\ &\quad \text{f. 更新 } m_i = m_{ij} \\ &4.\ \text{归一化 } O_i = O_i / \ell_i \end{aligned}

FlashAttention 实现细节

实现时通常先确定块大小,再设计数据搬运流水线,最后用编译器报告和性能分析工具检查寄存器压力。下面三个标签页对应这三个调优入口。

块大小(Block Size)选择

关键约束:共享内存容量限制

对于 MP31 架构(单 MP 192KB 共享内存):

SRAM 需求(Bc×d+2×Br×d)×sizeof(dtype)+额外状态缓冲\text{SRAM 需求} \approx (B_c \times d + 2 \times B_r \times d) \times \text{sizeof(dtype)} + \text{额外状态缓冲}

典型配置(d=128, FP16):

参数说明
BcB_c256Q 块大小
BrB_r128K/V 块大小
共享内存使用~192KB接近上限

不同隐藏维度(Head Dimension)

d推荐 Bc推荐 BrQ/K/V 基础 SRAM 使用
64512256约 128KB
128256128约 128KB
25612864约 128KB

实际共享内存使用还取决于双缓冲、online softmax 状态和其他临时缓冲,需要结合具体 kernel 实现确认。


FlashAttention-2/3 改进

FlashAttention-2 改进:序列并行

V1 问题

  • 单 Block 只计算 Q 的一个块
  • 长序列并行度不足

FlashAttention-2 改进

  • 在查询序列长度(Seq_Len_Q)上增加并行
  • 不同 ThreadBlock 并行处理不同 Q 分块,每个 ThreadBlock 只需加载一次对应的 Q 块
  • 内循环按键值序列长度(Seq_Len_KV)多次加载 K,V

显存与性能收益趋势

理论显存占用对比(FP16)

序列长度标准 AttentionFlashAttention降低倍数
N=5124MB0.5MB8x
N=204868MB2MB34x
N=4096268MB4MB67x
N=163844.3GB16MB268x

性能收益趋势(FP16)

序列长度标准 AttentionFlashAttention说明
N=512访存开销较低收益有限短序列下 kernel 开销占比更高
N=2048注意力矩阵访存压力增大收益明显分块和重计算开始体现优势
N=4096O(N²) 显存与访存压力明显收益显著避免落地完整注意力矩阵
N=16384显存与访存压力很高更适合长序列实际性能需以目标硬件实测为准

上表用于说明随序列长度增长的收益趋势,不代表固定硬件或固定 kernel 的性能承诺。


优化检查清单

块大小(Block Size)调优

  • Bc, Br 是否匹配 SRAM 容量?

    • 共享内存使用 < 90%
    • 留出余量给临时缓冲和 online softmax 状态
  • 是否考虑隐藏维度(Head Dimension)

    • d=64: 更大的 Bc, Br
    • d=256: 减小 Bc, Br

流水线优化

  • 是否使用双缓冲?

    • 两组共享内存交替使用
    • 加载与计算重叠
  • 张量内存引擎(TME)

    • 异步数据搬运
    • 减少寄存器压力

寄存器管理

  • 是否溢出?

    • 检查编译器报告的寄存器使用
    • 避免 spill 到局部内存
  • 占用率是否合理?

    • 使用 Moore Perf Compute 分析
    • 平衡占用率与寄存器使用

常见问题

Q1:FlashAttention 适合什么场景?

  • 更适合长序列和大 Batch 训练等注意力矩阵访存压力较高的场景。
  • 短序列下 kernel 启动、调度和同步开销占比更高,是否使用 FlashAttention 需要结合实测判断。
Q2:为什么 FlashAttention 更快?

  1. 减少全局内存访问(IO 感知)
  2. 避免存储 O(N²) 中间矩阵
  3. 更好的数据局部性
Q3:精度是否有损失?

  • FlashAttention 数学上等价于标准 Attention
  • 建议使用 FP32 累加器保证精度

相关文档