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 内存层次数据流 | 带宽利用率提升 |