数据截至 (上游 commit 0251105a2fb1)
02 · online softmax 的数学与实现
这一章讲什么: tiling(第 01 章)留下的那个拦路虎——softmax 分块后怎么算还对。先推数学,再逐行对照
csrc/flash_attn/src/softmax.h的真实 kernel 代码。这是整个 FlashAttention 的算法心脏。
1. 它要解决的小问题
safe softmax 的定义(减最大值防溢出):
m = max(x₁…x_N); p_i = e^(x_i − m) / Σ_j e^(x_j − m)
它需要两遍扫过整行:一遍求 max 和 sum,一遍归一化。
但 tiling 之后,每个 K 块只给你整行的一小段:处理第 1 块时,第 2 块里可能藏着更大的值——它会让已算的 max、sum、乃至已累加的输出全部作废。
问题一句话:能不能只往前扫一遍、不回看,还能在最后得到和整行 softmax 完全一样的结果?
2. 思路:不存结论,只存「能修正的中间态」
直觉分两步。
第一步:把「除法」推迟到最后。 扫描过程中不维护概率,只维护未归一化的量:
m_j:前 j 块的 running max;l_j:前 j 块以m_j为基准的指数和,l_j = Σ_{i≤j} e^(x_i − m_j);acc_o:前 j 块以m_j为基准的加权和,acc_o = Σ_{i≤j} e^(x_i − m_j)·V_i。
第二步:新块到来时,把旧账「折算」到新基准。 若新块的值让 max 从 m_old 涨到 m_new,旧账全部乘一个修正因子:
e^(x − m_old) = e^(x − m_new) · e^(m_new − m_old) ⇒ 旧量 × e^(m_old − m_new)
于是更新规则(这就是 online softmax 的全部):
m_new = max(m_old, rowmax(S_块))
l = l · e^(m_old − m_new) + rowsum(e^(S_块 − m_new))
acc_o = acc_o · e^(m_old − m_new) + e^(S_块 − m_new) · V_块
扫完所有块,最后才除:O = acc_o / l。数学上可以归纳证明:任意时刻 l 和 acc_o 都等于「以当前 m 为基准」的精确值,所以最终 acc_o / l 与两遍法逐位等价(浮点误差量级内)。
3. 一个数字例子
一行分数 x = [1, 3 | 2, 5],竖线是分块边界:
| 时刻 | m | l(未归一指数和) | 说明 |
|---|---|---|---|
| 处理块 1 后 | 3 | e^(1−3)+e^(3−3) = 1.135 |