跳到主要内容

数据截至 (上游 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。数学上可以归纳证明:任意时刻 lacc_o 都等于「以当前 m 为基准」的精确值,所以最终 acc_o / l 与两遍法逐位等价(浮点误差量级内)。


3. 一个数字例子

一行分数 x = [1, 3 | 2, 5],竖线是分块边界:

时刻ml(未归一指数和)说明
处理块 1 后3e^(1−3)+e^(3−3) = 1.135以 3 为基准
块 2 到来,新 max=551.135·e^(3−5) + e^(2−5)+e^(5−5) = 0.154+0.050+1旧账乘 e^(−2) 折算
收尾p = [e^(1−5), e^(3−5), e^(2−5), e^(5−5)] / l与整行 safe softmax 相同

重点:全程没有回头重读块 1。这就是「流过即算」。


4. 原理演示

# 示意,非源码 —— online softmax 的递推
def online_softmax_row(blocks):
m, l = -inf, 0.0 # running max、running sum
acc = 0.0 # 未归一化的加权和 Σ e^(x−m)·V
for S_blk, V_blk in blocks:
m_new = max(m, rowmax(S_blk))
corr = exp(m - m_new) # 旧账修正因子;首块 m=-inf 时 corr=0
P_blk = exp(S_blk - m_new)
l = l * corr + rowsum(P_blk)
acc = acc * corr + P_blk @ V_blk
m = m_new
return acc / l # 除法只做一次

对照第 1 节的更新规则,一行不多。真实 kernel 干的就是这件事——只是 S_blkacc 在寄存器里,exp 换成了 exp2,还叠了一串数值防护。


5. 真实实现:softmax.h 逐块对照

5.1 状态与更新规则

kernel 里每个线程持有 kNRows 行的两个标量,定义在 struct Softmax(csrc/flash_attn/src/softmax.h:129):

// csrc/flash_attn/src/softmax.h:129-134(节选,字段即 running 状态)
template <int kNRows>
struct Softmax {
TensorT row_max, row_sum; // 就是上面的 m 和 l

更新规则在 softmax_rescale_o(csrc/flash_attn/src/softmax.h:137)。非首块的分支做的事与第 2 节一一对应:

  • 先备份旧 max,再 reduce_max 出新 max(csrc/flash_attn/src/softmax.h:146-148);
  • 算修正因子并折算旧账:scores_scale = exp2f((scores_max_prev - scores_max_cur) * softmax_scale_log2),然后 row_sum *= scores_scaleacc_o_rowcol *= scores_scale(csrc/flash_attn/src/softmax.h:150-157);
  • 新块取指数再累加:scale_apply_exp2 + reduce_sum(csrc/flash_attn/src/softmax.h:162-165)。

被调用的位置正是前向主循环:带 mask 段在 csrc/flash_attn/src/flash_fwd_kernel.h:350,无 mask 段在 csrc/flash_attn/src/flash_fwd_kernel.h:414

5.2 exp2 与 FMA:白赚的性能

scale_apply_exp2(csrc/flash_attn/src/softmax.h:67)里有段注释,把动机写得很直白(csrc/flash_attn/src/softmax.h:79-81):不算 exp(x − max),改算 exp2(x·log₂e − max·log₂e)——这样编译器能把「乘 scale 再减」合成一条 FMA(乘加)指令,而不是分开的 fmul + fadd。配套地,softmax 的 scale 从 Python 侧一路带下来的就是 scale·log2(e),存在 Flash_fwd_params::scale_softmax_log2(csrc/flash_attn/src/flash.h:70)。exp2 本身在 GPU 上也是更基本的指令。

5.3 数值防护:处处防 NaN

  • 全被 mask 的行:max 是 −∞,−∞ − (−∞) = NaN。代码把 max_scaled 显式置 0(csrc/flash_attn/src/softmax.h:73-74 注释 + 代码),softmax_rescale_o 里也有同款 Check_inf 防护(csrc/flash_attn/src/softmax.h:151-153)。
  • 行和为 0 / NaN 的行:收尾时 inv_sum 置 1,LSE 置 ±∞(csrc/flash_attn/src/softmax.h:176-180)——因果掩码下「整行被 mask」的查询输出为 0,而不是 NaN。
  • 反向重算 P 时直接用 LSE 当指数基准(csrc/flash_attn/src/flash_bwd_kernel.h:536),天然稳定。

5.4 一次省事的延迟归约

softmax_rescale_o 里有句注释:row_sum 的跨线程归约不在循环里做,推迟到最后归一化时一起做(csrc/flash_attn/src/softmax.h:163-164)。循环内只做线程内累加——省掉每块一次的 warp shuffle。

需要跨线程归约时,用的是 quad_allreduce_(csrc/flash_attn/src/softmax.h:37):MMA fragment 的每一行恰好摊在 4 个线程上,所以 Allreduce<4> 三次 shuffle 就完成行内 max/sum。

5.5 收尾:归一化 + 写出 LSE

normalize_softmax_lse(csrc/flash_attn/src/softmax.h:170)做两件事:

  1. acc_o1/l,得到最终输出(csrc/flash_attn/src/softmax.h:182);
  2. 返回 lse = row_max·softmax_scale + log(row_sum)(csrc/flash_attn/src/softmax.h:180)——这就是反向重算 P 要用的 logsumexp,前向 epilogue 把它写回 HBM(csrc/flash_attn/src/flash_fwd_kernel.h:440)。

6. 关键细节 / 坑

  1. 首块要特判。 Is_first=true 时旧账不存在,直接初始化而不是 rescale(csrc/flash_attn/src/softmax.h:141-144);示意代码里的 corr=0 对应的就是这件事。
  2. LSE 的符号惯例要小心。 行和为 0 时普通 kernel 写 +INFINITY,split-KV 分支写 -INFINITY(csrc/flash_attn/src/softmax.h:179),因为 combine kernel 把 −∞ 当「这个 split 没贡献」的标记(见第 03 章 split-KV 一节)。
  3. 修正因子的尺度也被预先乘进 log2 域。 exp2f((m_prev − m_cur) * softmax_scale_log2)(csrc/flash_attn/src/softmax.h:154)——所有指数运算统一在 2 的幂域,避免混用 e 底和 2 底。
  4. 这套数学不是本库首创。 online softmax 源自 Milakov & Gimelshein(2018)的在线归一化;本库的贡献是把它和 tiling、重计算、tensor core 布局缝成一个 kernel(inferred——论文 arXiv:2205.14135 有相关工作的完整脉络,入口见 README.md:8)。

7. 本章小结

  • online softmax = 维护 (m, l, acc_o) 三元组,新块到来用 e^(m_old − m_new) 折算旧账,除法推迟到末尾。
  • 代码里的每个数字技巧都有出处:exp2+FMA、−∞ 防护、延迟归约、LSE 输出。
  • 有了它,第 01 章的 tiling 骨架才成立;下一章看完整的 kernel 主循环怎么把它跑起来。