推理服务 · 内核

FlashAttention

精确注意力,在片上内存中分块计算

结果
精确 与标准实现相同
内存
O(N) 标准实现 O(N²)
A100 SRAM
192 KB × 108 个 SM
A100 HBM
40–80 GB
图中数值
玩具:8 个词元,d 4,分块 4 × 4

全部步骤

  1. 快内存与慢内存

    GPU 在芯片上有少量非常快的内存(SRAM,紧挨着运算单元),旁边还有大量较慢的内存(HBM)。在两者之间搬运数据往往比算术本身更耗时,所以快速的内核会把工作留在 SRAM 里。

    SRAM 19 TB/s · HBM 1.5 TB/s

  2. 标准注意力

    普通注意力分几个独立的步骤运行,每一步都从 HBM 读输入,再把结果写回。其中两个结果,即分数 S 和权重 P,是 N × N 的:输入很长时,它们在流量和内存中都占大头。

    S、P:N × N,写入 HBM 再读回

  3. 在线 softmax

    softmax 要先知道一行的最大值和总和,任何权重才能定下来。在线 softmax 维护一个滑动最大值 m 和总和 ℓ:当新来的块带来更大的最大值时,之前的一切都乘以 e^(mold − mnew) 重新缩放。结果与普通 softmax 完全相同。

    ℓ ← e^(mold − mnew) ℓ + Σ e^(s − mnew)

  4. 逐块计算

    FlashAttention 加载一块查询,然后让键和值的块依次流过 SRAM。每个 4 × 4 的分数分块只存在于芯片上;每来一个分块,就重新缩放滑动的 m、ℓ 和输出,最后除以 ℓ。真实的算术,已与普通注意力核对。

    Oi = Σj e^(Sij − m) Vj / ℓ,逐块计算

  5. 更少的数据搬运,没有 N × N

    N × N 的矩阵从不存储,所以注意力的内存随 N 而不是 N² 增长,在 HBM 和芯片之间传输的数据也少得多。论文报告注意力快了 2 到 4 倍,结果完全一致;反向传播时重算分块,而不是把它们存下来。

    内存 O(N) · 流量约低 4 倍

代码

# standard: three kernels, S and P (N × N) go through HBM
S = Q @ K.T / sqrt(d); P = softmax(S); O = P @ V
# FlashAttention: one kernel; for each block of queries, stream K and V
for i in query_blocks:                          # Q_i stays in SRAM
    m, l, acc = -inf, 0, 0
    for j in key_blocks:                        # K_j, V_j loaded into SRAM
        s = Q[i] @ K[j].T / sqrt(d)             # one tile, never written out
        m_new = max(m, s.max(-1))
        p = exp(s - m_new)
        l = exp(m - m_new) * l + p.sum(-1)
        acc = exp(m - m_new) * acc + p @ V[j]
        m = m_new
    O[i] = acc / l                              # exactly softmax(S) @ V

延伸阅读