CUDA 算子工程:手写 FlashAttention v2 之路
第 6 章 Softmax 与 Online Softmax
Online softmax 不是 FlashAttention 发明的:它出自 Milakov & Gimelshein 的 Online normalizer calculation for softmax(NVIDIA,2018,arXiv:1805.02867)。 FlashAttention 做的事,是把这套一遍式的归一化和 attention 的分块矩阵乘缝到了一起。
6.1 Softmax 的标准定义
给定一个向量 ,softmax 定义为:
直接照定义计算需要两遍:
// 第一遍: 算分母 (求和)
float sum = 0;
for (int i = 0; i < N; ++i) sum += expf(x[i]);
// 第二遍: 算每个输出
for (int i = 0; i < N; ++i) y[i] = expf(x[i]) / sum;
但这段代码有一个致命问题:数值溢出。
6.2 数值稳定的 Safe Softmax
如果 中有比较大的数(比如 100),那 大约是 。FP32 的上限大约是 ——直接溢出成 +inf。如果 全是 1000, 直接溢出,所有项都变 inf,分子分母 inf/inf 输出 NaN。
LLM 训练和推理中,attention 的 logits()没有取值上界的保证。若在 FP16 里算 exp,风险更高:FP16 上限只有 65504, 就溢出。BF16 的指数位与 FP32 同为 8 位,表示范围和 FP32 相当,但尾数只有 7 位,所以实践中 softmax 一般在 FP32 里算。
解决方法是经典的 safe softmax 技巧——同时减去最大值:
数学上完全等价(分子分母都乘以 抵消),但所有 ,绝不溢出。
但 safe softmax 现在需要 3 遍:
// Pass 1: 找 max
float m = -INFINITY;
for (int i = 0; i < N; ++i) m = fmaxf(m, x[i]);
// Pass 2: 算分母
float sum = 0;
for (int i = 0; i < N; ++i) sum += expf(x[i] - m);
// Pass 3: 算每个输出
for (int i = 0; i < N; ++i) y[i] = expf(x[i] - m) / sum;
3 遍意味着 x 数组要从 HBM 读 3 次(一行大到缓存装不下时)。这是大问题——softmax 的算术强度本身就低(每个元素几次浮点操作 vs 4 字节读),读 3 遍再加写回 ,每个元素 4 次访存,而下限是 1 读 1 写共 2 次。
6.3 Online Softmax:1 遍数学
Online Softmax 的核心思想:在遍历数据的过程中同时维护当前的 max 和 sum,不需要预先知道全局 max。
6.3.1 数学推导
定义两个状态量:
- = 截至第 个元素时的最大值
- = 截至第 个元素的"调整后 exp 和"
初始状态 ,。
当看到新元素 时,更新规则:
关键的"修正项"是 :当 max 被更新( 时这个值小于 1),把之前累积的 sum 也"按比例缩小",保证它仍然是相对于新 max 的 sum。
把这个递推走完, 就是 safe softmax 的分母(原论文把它记作 ,本专栏沿用 FlashAttention 论文和第 14–16 章的记号 )。然后再走一遍计算输出(这一遍可以和下游计算 fuse 在一起)。
如果是 attention 这种"sum 之后还要点积"的场景,可以做到真正的 1 pass——这是 FA 的核心。
6.3.2 验证一个简单例子
考虑 。
朴素方式:
Online:
| 步骤 | |||
|---|---|---|---|
| 初始 | - | 0 | |
| n=1 | 1 | 1 | |
| n=2 | 5 | 5 | |
| n=3 | 3 | 5 |
,和朴素方式得到的 完全一致。
6.3.3 推导的几何理解
为什么 是正确的修正?
设 。我们希望 。
把它拆开:
第二项可以改写:
所以:
干净的代数变换。这就是为什么 online softmax 数学上严格等价于 safe softmax——它只是把"先扫一遍找 max 再扫一遍累加"改写成"边扫边更新"。
6.4 Online Softmax 的 GPU 实现
把 online softmax 实现成 GPU kernel。基本模板:
__global__ void softmax_online(const float* in, float* out, int N) {
int tid = threadIdx.x;
extern __shared__ float smem[];
// 约定: 一个 block 处理一行, blockDim.x == 256 (8 个 warp),
// 启动时动态 SMEM 至少 16 个 float (64 字节)
// ====== Phase 1: Online sweep (1 pass) ======
float m = -INFINITY; // 当前 max
float l = 0.0f; // 当前调整后 sum
for (int i = tid; i < N; i += blockDim.x) {
float x = in[i];
float m_new = fmaxf(m, x);
if (m_new == -INFINITY) continue; // 目前为止全是 -INF: 否则 expf(-INF - (-INF)) = NaN
l = l * expf(m - m_new) + expf(x - m_new);
m = m_new;
}
// 此时每个 thread 持有局部的 (m, l)
// ====== Phase 2: Block-level reduce ======
// Warp 内合并 (m, l)
auto warp_combine = [](float& m, float& l, float m2, float l2) {
float m_new = fmaxf(m, m2);
if (m_new == -INFINITY) return; // 两边都是空状态 (-INF, 0), 保持不动
l = l * expf(m - m_new) + l2 * expf(m2 - m_new);
m = m_new;
};
for (int offset = 16; offset > 0; offset >>= 1) {
float m2 = __shfl_xor_sync(0xFFFFFFFF, m, offset);
float l2 = __shfl_xor_sync(0xFFFFFFFF, l, offset);
warp_combine(m, l, m2, l2);
}
// 此时 warp 内所有 lane 都持有 warp 的 (m, l)
// 写入 SMEM, block 内再 reduce 一次
int warp_id = tid / 32;
int lane_id = tid % 32;
if (lane_id == 0) {
smem[warp_id * 2 + 0] = m;
smem[warp_id * 2 + 1] = l;
}
__syncthreads();
// Warp 0 收集 8 个 warp 的 (m, l)
if (warp_id == 0) {
m = (lane_id < 8) ? smem[lane_id * 2 + 0] : -INFINITY;
l = (lane_id < 8) ? smem[lane_id * 2 + 1] : 0.0f;
// 注意 mask 必须是 0xFFFFFFFF 而不是 0xFF:这里的守卫是
// `warp_id == 0`,warp 0 的 32 个 lane **全都**在执行这条指令,
// 而 `*_sync` 系列要求"执行它的线程必须被 mask 点名",
// 写 0xFF 会让 lane 8..31 落在 mask 之外,行为未定义。
// lane 8..31 持的是 (-INF, 0) 这个单位元(combine 里的
// -INF 判断保证它真是单位元),跟着算也不会污染结果。
for (int offset = 4; offset > 0; offset >>= 1) {
float m2 = __shfl_xor_sync(0xFFFFFFFF, m, offset);
float l2 = __shfl_xor_sync(0xFFFFFFFF, l, offset);
warp_combine(m, l, m2, l2);
}
if (lane_id == 0) {
smem[0] = m;
smem[1] = l;
}
}
__syncthreads();
float final_m = smem[0];
float final_l = smem[1];
// ====== Phase 3: Normalize ======
for (int i = tid; i < N; i += blockDim.x) {
out[i] = expf(in[i] - final_m) / final_l;
}
}
关键点:
- Phase 1 每个 thread 维护自己的
(m, l)状态。 - Phase 2 用
warp_combine合并不同 thread 的(m, l)——这是 online softmax 的核心组合规则。 - Phase 3 用最终的全局
(m, l)做归一化。 - 两处
m_new == -INFINITY判断不能省:线程分到的元素全是-INF(被 mask 掉),或者 N 小于 blockDim.x、有线程一个元素都没分到,它的状态就停在 ;两个这样的状态相遇时expf(-INF - (-INF))是expf(NaN),NaN 会一路传进最终结果。用 numpy 按本 kernel 的线程划分和蝶形归约逐步模拟,去掉这两行时 N=100(< 256)或前 300 个元素为-INF的行输出全是 NaN,加上后与torch.softmax的最大误差在 量级。FA2 用的是同一类防护:csrc/flash_attn/src/softmax.h在 row max 为-INFINITY时改用 0 去算缩放因子。
本 kernel 用 nvcc 13.4 以 -arch=sm_90a 编译通过,21 个寄存器、无 spill(nvcc 13.4 编译所得,不同 CUDA 版本可能不同)。
这里 Phase 1 和 Phase 3 都需要读 in[i]——所以严格说还是 2 pass。但 Phase 1 不再需要 max、sum 两次单独的 pass——它把 max 和 sum 合并到一遍里了。这对 SRAM 受限的 GEMM/attention 场景非常关键,因为可以让 in 数据只在 SMEM/寄存器里待一次。
6.4.1 Combine 函数的对称性
注意 warp_combine 函数:
auto warp_combine = [](float& m, float& l, float m2, float l2) {
float m_new = fmaxf(m, m2);
if (m_new == -INFINITY) return; // 两边都是空状态 (-INF, 0), 保持不动
l = l * expf(m - m_new) + l2 * expf(m2 - m_new);
m = m_new;
};
这个函数是对称、可结合的:
证明只需一步:把状态 看作它代表的量 ,combine 就是把两者代表的量相加,再改用新的 max 作基准表示。无论先合并哪两个, 合并的结果都是 ,其中 ;交换律从公式的对称性直接可见。原论文(式 4)给出了同一个算子,并声明它可结合、可交换,证明从略。注意这是实数意义上的等式:浮点下不同的合并顺序会带来末位舍入差异,结果不保证逐位相同。
结合律是并行归约的必要条件——没有它,改变合并顺序就可能改变结果;交换律则保证上面 __shfl_xor_sync 蝶形归约里配对的两个 lane 以相反顺序合并,也得到同一个值。再加上单位元 (前提是上面那个 -INFINITY 判断),combine 构成一个幺半群(monoid),这正是 online softmax 能套进任何并行归约框架的原因。
6.5 LSE 形式:Log-Sum-Exp
attention 中常用 softmax 的 log 版本(log-softmax),定义为:
其中 是 log-sum-exp。
数值稳定的 LSE:
在 online 形式中,存 时实际上已经在算 LSE:
FA2 反向传播时需要存 LSE(一个标量/行)来支持梯度计算。FA2 前向除了输出 O,只额外存下 LSE,不存 的 attention 矩阵——这是它额外显存从 降到 的关键。
6.6 Online Softmax 应用到 FA:1 Pass 真的成立
FA 的核心创新就是把 attention 改写成 online 形式:
朴素 attention:
需要 3 个完整的 pass:
- 算 (HBM 写 N×N 矩阵)。
- 算 行级(HBM 写 N×N 矩阵)。
- 算 (HBM 写 N×d 矩阵)。
中间结果 和 都是 大小:N=4096 时单个 head 就是 个元素,FP16 下 32 MB;乘上 head 数和 batch,写一次读一次的 HBM 流量非常可观。
FA 用 online softmax 把它压成 1 pass:
# 伪代码: 单个 query 行 q, K/V 按 B 行一块
m, l = -inf, 0
O = zeros(d) # output accumulator
for j in range(0, N, B): # block-by-block
K_block = K[j:j+B]
V_block = V[j:j+B]
S_block = q @ K_block.T # B 个分数
m_block = max(S_block) # block 局部 max
P_block = exp(S_block - m_block) # 局部 exp
l_block = sum(P_block)
# Combine 到 (m, l)
m_new = max(m, m_block)
alpha = exp(m - m_new)
beta = exp(m_block - m_new)
# 重要: 之前累积的 O 也要按 alpha 缩放
O = O * alpha + beta * (P_block @ V_block)
l = l * alpha + beta * l_block
m = m_new
O = O / l # 最终归一化
注意几点:
- m, l, O 三个状态量同步演化。每来一个 K/V block,都更新这三个。
- O 的缩放因子 来自 online softmax 的修正——之前累积的 attention 输出也要按新 max 缩放。
- 完全不写中间 S 或 P 矩阵到 HBM——只在寄存器/SMEM 中流过。
- 这里沿用了 6.4 节 combine 的写法:块内先减局部 max,再用 对齐到新 max。第 14、15 章的 FA2 写法更省一步:直接用 算 ,于是不再需要 ,两者数学上等价(用 numpy 对单个 query 行分块复算,与
softmax(qK^T)V的差在 量级)。
这就是 FA 的本质:用 online softmax 让 attention 变成可流式的算法。Q/K/V 可以按 tile 喂给 kernel,算完就丢,不需要存中间结果。
第 14-15 章会把这个伪代码落到具体的 CUDA kernel 上。
6.7 Softmax 的访存优化
回到 standalone softmax 的 kernel 实现。性能优化要点:
6.7.1 整行处理 vs 分块处理
LLM 中 softmax 是按行操作的(attention 的每一行独立 softmax)。两种 block 配置:
整行处理:一个 block 处理一行的所有 N 个元素。
适合: 行长适中 (一个 block 就能处理完, 最好整行放得进 SMEM 或寄存器)
优势: 不需要跨 block 通信
分块处理:多个 block 协作处理一行。
适合: N 极大 (比如长序列 attention 的中间矩阵)
代价: 需要两阶段 + atomic 或 cluster reduce
一行交给一个 block 能处理完时,整行处理更简单:不需要跨 block 通信,也不需要第二个 kernel 或 atomic。
6.7.2 Vectorized I/O
和 reduce 类似,softmax 的访存也应该 vectorized:
// in 为 const half*; 每线程一次读 4 个 fp16 = 8 字节 (i 须是 4 的倍数, 地址 8 字节对齐)
uint2 packed = *reinterpret_cast<const uint2*>(&in[i]);
half2 v0 = *reinterpret_cast<half2*>(&packed.x);
half2 v1 = *reinterpret_cast<half2*>(&packed.y);
别写成两次 half2 读:*reinterpret_cast<const half2*>(&in[i]) 和 &in[i + 2] 各读 4 字节,编译器并不会替你合并——nvcc 13.4 在 sm_90a 上编出的是两条 32 位的 LDG.E,而上面的 uint2 写法是一条 LDG.E.64(nvcc 13.4 编译所得,不同 CUDA 版本可能不同)。
或者一次读 16 字节,用 int4 装 8 个 fp16(这时每个线程负责下标 i * 8 起的 8 个元素,地址须 16 字节对齐,nvcc 13.4 编出一条 LDG.E.128):
int4 packed = *reinterpret_cast<const int4*>(&in[i * 8]);
half h[8];
memcpy(h, &packed, 16);
6.7.3 Fused Softmax + Mask
attention 中 softmax 之前通常有 mask(causal mask、padding mask)。fused 写法:
float x = in[i] + mask[i]; // mask 通常是 0 或 -INFINITY
m = fmaxf(m, x);
把 mask add 直接 fuse 到 softmax 的 max/sum 阶段,避免单独走一遍 mask kernel。用 -INFINITY 做 mask 时,online 更新里 6.4 节那个 m_new == -INFINITY 判断就是必需的,否则被整段 mask 掉的线程会产生 NaN。整行都被 mask 时,结果本来就没有定义:torch.softmax 对全 -inf 的行输出 NaN,加了判断的 kernel 也输出 NaN()。
6.7.4 三种写法的访存账
原论文按"每个元素访问几次内存"给三种写法记账(Milakov & Gimelshein 2018,第 2–4 节):
| 实现 | 每元素访存次数 | 构成 |
|---|---|---|
| 朴素 softmax(不减 max,不安全) | 3 | 读 2 遍 + 写 1 遍 |
| Safe softmax(3-pass) | 4 | 读 3 遍 + 写 1 遍 |
| Online softmax | 3 | 读 2 遍(max 与 sum 合成一遍)+ 写 1 遍 |
Online 相对 safe 省下的访存是 倍。论文在 Tesla V100(fp32,每批 4000 个向量)上实测:向量长度到 1000 左右时三种写法差不多(数据还在 L1/L2 里),之后被 DRAM 带宽卡住,online 相对 safe 很快达到约 1.3 倍,与 1.33 倍的访存比吻合。对 standalone softmax 来说,这差不多就是这项改写的上限——它本来就是带宽 bound 的算子,每元素至少要 1 读 1 写。论文里真正大的收益来自融合:Softmax + TopK 融合后每元素只剩 1 次访存,实测最高约 5 倍。FA 把 softmax 嵌进 attention 也是同一个道理,收益来自根本不写 、 这两个 中间矩阵,那是另一个量级的事。
6.8 这一章给我们的"内核数学"
Online softmax 是这个专栏第一次正式接触"用算法重写让 GPU 友好"的思路。它带给读者两个核心 insight:
-
数学等价不等于 GPU 等价。同一个 softmax 公式有 3 遍和 2 遍两种数学等价的算法(与下游融合后还能压到 1 遍),访存次数不同,GPU 上的性能也就不同:standalone 时约 1.3 倍,融合进下游后差距可以到数倍。算法重写和"低层 kernel 优化"是 GPU 性能工程的两个独立维度,online softmax 是前者的典范。
-
Streaming(流式)算法在 GPU 上是黄金。能够用单 pass 维护少量状态量来累积结果的算法,天然适合 GPU——因为它把"中间结果"留在寄存器/SMEM,避免 HBM 往返。Online softmax、Welford 在线方差(下一章)、prefix sum——这些都是流式算法的代表。
第 7 章我们把 online 思路用到另一个 LLM 高频算子上:LayerNorm 与 RMSNorm。LayerNorm 需要算均值和方差——传统写法要么先算均值再算方差、多读一遍,要么用 一遍算完却有数值稳定问题,Welford 算法(online softmax 论文自述,其灵感正来自 Welford 的数值稳定在线方差算法)可以一遍同时算出均值和方差。读完第 7 章,读者会发现"online 思维"在 LLM 算子里几乎无处不在。
本章动手练习:
- 实现一个 N=4096 的整行 softmax kernel,先用 3-pass,再改成 online,对比性能。
- 用 PyTorch 的
torch.nn.functional.softmax跑一遍,用 Nsight Compute 看它实际落到了哪个 kernel 上——PyTorch v2.11 对最后一维做 softmax 时,短行(dim_size <= 2048且整行不超过 8 KB,见pytorch-v2.11.0/aten/src/ATen/native/cuda/SoftMax.cu:1097)走pytorch-v2.11.0/aten/src/ATen/native/cuda/PersistentSoftmax.cuh的softmax_warp_forward(:68);更长的行(比如 N=4096)在pytorch-v2.11.0/aten/src/ATen/native/cuda/SoftMax.cu的五个前向 kernel 里选:cunn_SoftMaxForwardFast(:702)、cunn_SoftMaxForward(:734)、cunn_SoftMaxForwardReg(:770)、cunn_SoftMaxForwardGmem(:827)、cunn_SoftMaxForwardSmem(:888)。分派逻辑在同文件 :1108 起:use_fast_softmax为真时按dim_size % ILP在 Gmem / Fast 之间选;否则先算potential_register_count决定能不能整行塞进寄存器(Reg),塞不下就看can_use_smem(同时要求整行放得进sharedMemPerBlock、输入输出都 16 字节对齐、dim_size是 ILP 的整数倍)走 Smem,再不行才退回最通用的cunn_SoftMaxForward。- 思考:如果 softmax 的 N 极大(比如 N=1M),整行处理放不进单 block 的 SMEM,怎么用 online softmax + cluster reduce 解决?