Skip to content
🔗 分享本题
查看我的学习进度 →

FlashAttention 通过片上分块和在线 softmax 减少 HBM 读写的机制图

🧠 记忆锚点:加速来自 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 读 HBM

FlashAttention 优化:

核心思想: 分块计算,避免物化(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

内存访问对比:

操作传统 AttentionFlashAttention
HBM 读4n²d4nd
HBM 写2n²d2nd
总访问O(n²d)O(nd)
加速比1xn/d

实测性能(A100 GPU):

序列长度传统 AttentionFlashAttention加速比
51210ms8ms1.25x
2K150ms50ms3x
8K2.4s400ms6x
128KOOM10s-

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(原论文)