
🧠 记忆锚点:加速来自 IO-aware 分块与在线 softmax,不是近似注意力;少写 HBM 才是关键。
💡 答案要点
FlashAttention = 优化 Attention 计算的 I/O 效率
传统 Attention 的问题:
传统 Attention 计算(三步):
1. S = QK^T(n×d @ d×n = n×n)
2. P = softmax(S)(n×n)
3. O = PV(n×n @ n×d = n×d)
问题:
中间矩阵 S, P 的大小是 n×n
- n=4K:16M 元素
- n=128K:16B 元素(显存炸了)
需要多次读写 HBM(高带宽内存):
QK^T 写 HBM → softmax 读 HBM → PV 读 HBMFlashAttention 优化:
核心思想: 分块计算,避免物化(materialize)大矩阵
1. 将 Q, K, V 分块(tile)
2. 每个块加载到 SRAM(片上内存)
3. 在 SRAM 内完成计算
4. 只写回最终结果
好处:
- 减少 HBM 访问(慢)
- 增加 SRAM 访问(快 10x)
- 不需要存储 n×n 的中间矩阵算法流程:
python
# 传统(朴素实现)
S = Q @ K.T # 写 HBM(n×n)
P = softmax(S) # 读+写 HBM
O = P @ V # 读 HBM
# FlashAttention
block_size = 128
for i in range(0, n, block_size):
# 加载块到 SRAM
Q_block = load_to_sram(Q[i:i+block_size])
for j in range(0, n, block_size):
K_block = load_to_sram(K[j:j+block_size])
V_block = load_to_sram(V[j:j+block_size])
# 在 SRAM 内完成计算
S_block = Q_block @ K_block.T
P_block = softmax(S_block)
O_block += P_block @ V_block
# 写回 HBM
O[i:i+block_size] = O_block内存访问对比:
| 操作 | 传统 Attention | FlashAttention |
|---|---|---|
| HBM 读 | 4n²d | 4nd |
| HBM 写 | 2n²d | 2nd |
| 总访问 | O(n²d) | O(nd) |
| 加速比 | 1x | n/d 倍 |
实测性能(A100 GPU):
| 序列长度 | 传统 Attention | FlashAttention | 加速比 |
|---|---|---|---|
| 512 | 10ms | 8ms | 1.25x |
| 2K | 150ms | 50ms | 3x |
| 8K | 2.4s | 400ms | 6x |
| 128K | OOM | 10s | - |
FlashAttention-2 改进:
1. 减少非矩阵乘法运算(softmax 优化)
2. 更好的并行化
3. 支持更长序列(256K+)
4. 速度再提升 2x面试话术:
"FlashAttention 的核心是 I/O 优化。传统 Attention 要读写 n² 的中间矩阵,FlashAttention 分块计算避免了物化。在长序列(8K+)上加速 5-10 倍,而且支持更长上下文。"
📚 参考:FlashAttention: Fast and Memory-Efficient Exact Attention(原论文)