推理服务 · 内核
- 结果
- 精确 与标准实现相同
- 内存
- O(N) 标准实现 O(N²)
- A100 SRAM
- 192 KB × 108 个 SM
- A100 HBM
- 40–80 GB
- 图中数值
- 玩具:8 个词元,d 4,分块 4 × 4
全部步骤
快内存与慢内存
GPU 在芯片上有少量非常快的内存(SRAM,紧挨着运算单元),旁边还有大量较慢的内存(HBM)。在两者之间搬运数据往往比算术本身更耗时,所以快速的内核会把工作留在 SRAM 里。
SRAM 19 TB/s · HBM 1.5 TB/s标准注意力
普通注意力分几个独立的步骤运行,每一步都从 HBM 读输入,再把结果写回。其中两个结果,即分数 S 和权重 P,是 N × N 的:输入很长时,它们在流量和内存中都占大头。
S、P:N × N,写入 HBM 再读回在线 softmax
softmax 要先知道一行的最大值和总和,任何权重才能定下来。在线 softmax 维护一个滑动最大值 m 和总和 ℓ:当新来的块带来更大的最大值时,之前的一切都乘以 e^(mold − mnew) 重新缩放。结果与普通 softmax 完全相同。
ℓ ← e^(mold − mnew) ℓ + Σ e^(s − mnew)逐块计算
FlashAttention 加载一块查询,然后让键和值的块依次流过 SRAM。每个 4 × 4 的分数分块只存在于芯片上;每来一个分块,就重新缩放滑动的 m、ℓ 和输出,最后除以 ℓ。真实的算术,已与普通注意力核对。
Oi = Σj e^(Sij − m) Vj / ℓ,逐块计算更少的数据搬运,没有 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