返回首页

Softmax / LayerNorm 融合

Lesson 09 · CUDA 性能调优 · 把 reduction 与逐元素算子合成一个 kernel

上一节把 reduction 调到了带宽极限——单个求和算子,我们亲手验证了 "天花板是带宽不是算力",且 grid-stride 之后再精修归约树收益≈0。这一节回答那节课结尾留的问题: 既然单个算子已经吃满带宽,再快从哪来? 答案是融合(fusion)——把多个算子合成一个 kernel, 省掉中间结果在 HBM 上的往返。Softmax 和 LayerNorm 是 AI 里最典型的"reduction + 逐元素"组合,正是融合的主战场。

本节目标(一个可带走的胜利): 你能说清"未融合的 softmax 为什么要把数据读 3–5 遍", 能手写一个每行一个 block、一趟读完的融合 softmax kernel(online 算法), 并在 3060 上实测它比"多 kernel 拼接"版快多少——再用带宽算出它离极限还差几成。

1. 先看未融合版:数据被读了几遍?

Softmax 对一行 x[0..N) 做:y_i = exp(x_i - max) / Σexp(x_j - max)。 减 max 是为了数值稳定(防 exp 溢出)。如果你用现成算子库逐个拼,典型是 4 个 kernel:

// 朴素四趟:每一趟都把整行从 HBM 读出 / 写回一次
m   = reduce_max(x)        // 趟1:读 x
e   = exp(x - m)           // 趟2:读 x,写 e
s   = reduce_sum(e)        // 趟3:读 e
y   = e / s                // 趟4:读 e,写 y

数一下 HBM 流量:x 读 2 次、e 写 1 读 2、y 写 1 —— 对一行 N 个 float,光是访存就是 6N 字节量级(理想下限只需读 x 一遍、写 y 一遍 = 2N)。 而每个 kernel 还各有一次启动开销 + 把行数据反复进出 L2/HBM。reduction 那节告诉我们: 访存瓶颈算子的耗时正比于搬的字节数。读 3 倍的数据,就是慢 3 倍。

融合的本质 = 减少 HBM 往返。 把这 4 步塞进一个 kernel:一行数据进 shared/寄存器后, max、exp、sum、除法全在片上做完,只读一次 x、写一次 y。中间量 e 根本不落 HBM。 这就是 AI 算子提速的主线——不是把某一步算得更快,而是让数据少搬几趟

2. 融合 softmax:每行一个 block,一趟读完

布局:输入是 [rows × N],一个 block 负责一行。block 内线程用 grid-stride 把这一行扫一遍, 同时做 reduction。关键技巧叫 online softmax:一边读一边维护 (running_max, running_sum), 读到更大的值时把已累加的 sum 按比例缩放。这样只需扫一遍 x 就同时拿到 max 和 sum,而不是分两趟。[1]

// 每个线程先在寄存器里 online 归约自己负责的那些元素
float m = -INFINITY, s = 0.0f;
for (int i = tid; i < N; i += blockDim.x) {
    float x = row[i];
    float m_new = fmaxf(m, x);
    s = s * __expf(m - m_new) + __expf(x - m_new);  // 旧 sum 按新 max 缩放
    m = m_new;
}

每个线程算完自己那份 (m, s) 后,block 内还要把这些局部 (max,sum) 对合并成全行的一个。 合并规则和上面一行一样:两个 (m₁,s₁)(m₂,s₂)M=max(m₁,m₂), S=s₁·exp(m₁−M)+s₂·exp(m₂−M)。用上一节warp shuffle 做这个归约—— 只是被归约的不再是单个 float,而是这个二元组。

// warp 内对 (m,s) 二元组做 shuffle 归约
for (int off = 16; off > 0; off >>= 1) {
    float m2 = __shfl_down_sync(0xffffffff, m, off);
    float s2 = __shfl_down_sync(0xffffffff, s, off);
    float M = fmaxf(m, m2);
    s = s * __expf(m - M) + s2 * __expf(m2 - M);
    m = M;
}

拿到全行的 (M, S) 后(广播给所有线程),第二趟再扫一遍写出 y_i = exp(x_i − M) / S。注意这第二趟读 x 的代价无法避免—— 要写 y 就得再读一次 x(除非 N 小到能整行缓存进 shared,那样连第二趟读 HBM 都省了)。

online softmax 把"max 一趟、sum 一趟"压成一趟。 配合每行一个 block, 整个 softmax 只剩 读 x 两趟(归约一趟 + 写出一趟)、写 y 一趟,从 6N 砍到 3N。 若整行装得进 shared(N ≤ 几千),第二趟从 shared 读 → HBM 流量逼近理想的 2N。

3. N 装得下时:整行进 shared,逼近 2N 下限

当 N 不大(比如 transformer 的 vocab 之外的多数归一化,N=768/1024/4096), 一行能整个放进 shared memory。那么策略变成:第一趟把整行读进 shared(顺带 online 归约), 第二趟从 shared 读、算 exp、写 y。HBM 上 x 只读一次、y 只写一次 = 2N,触到理想下限

extern __shared__ float srow[];   // 动态 shared,大小 = N*4
// 趟1:读 HBM → shared,同时累积局部 (m,s)
for (int i = tid; i < N; i += blockDim.x) srow[i] = row[i];
// ...block 内归约出全行 (M,S)...
// 趟2:从 shared 读,无需再碰 HBM 的 x
for (int i = tid; i < N; i += blockDim.x)
    out[i] = __expf(srow[i] - M) / S;

这正是 Lesson 04 shared memory 的标准用法: 一份数据要被读多次时,先搬进片上,后续都从片上读。softmax 的 x 被读两次(归约 + 写出), 所以缓存它正好划算。

4. LayerNorm:同一套骨架,换成 (sum, sum²)

LayerNorm 对一行做:y = (x − μ) / √(σ² + ε) · γ + β,其中 μ 是均值、σ² 是方差。 结构和 softmax 一模一样——先一趟归约求统计量,再一趟逐元素归一化。区别只在归约的是什么:

SoftmaxLayerNorm
归约量(max, Σexp)(Σx, Σx²)
统计量M, Sμ=Σx/N, σ²=Σx²/N − μ²
逐元素exp(xᵢ−M)/S(xᵢ−μ)/√(σ²+ε)·γ+β
// LayerNorm 第一趟:一次 shuffle 归约同时拿到 sum 和 sum_sq
float sum = 0, sum_sq = 0;
for (int i = tid; i < N; i += blockDim.x) {
    float v = row[i];
    sum += v;  sum_sq += v * v;
}
// 对 (sum, sum_sq) 各做一次 warp shuffle + block 归约
float mean = sum / N;
float var  = sum_sq / N - mean * mean;   // 一遍法:E[x²]-E[x]²
一遍法方差 E[x²]−E[x]² 在 fp32 下可能因抵消而掉精度。 工业实现(如 Apex、cuDNN)多用 Welford 算法:online 维护 (count, mean, M2), 数值更稳,且天然适合 warp 归约。softmax 的 online 缩放、LayerNorm 的 Welford—— 本质都是"一趟拿到本需多趟的统计量"[2]

5. 实测:融合 vs 多 kernel(3060 / sm_86)

评判融合好不好,还是回到带宽。融合版的理想 HBM 流量是 (读x + 写y); 未融合版是它的 2–3 倍。所以预期加速比≈流量比,而不是算力比。编译跑(3060 用 sm_86,A100 用 sm_80):

nvcc -O3 -arch=sm_86 softmax_fused.cu -o sm
./sm                     # 比对 fused vs 4-kernel 的耗时与等效带宽
看什么含义期望
等效带宽 = 总字节/耗时融合版总字节≈2N·rows✅ 上到峰值 70%+
fused / 4-kernel 加速比≈ 流量比 ≈ 2–3x✅ 与字节比吻合
DRAM Throughput(ncu)融合后应接近 SOL Memory⚠️ WSL 下采不到,见下
WSL 提醒:你的 3060 在 WSL2 下,ncu 的硬件性能计数器不可用 (实测会报 "Profiling is not supported on device ... WSL")。所以这一节的验证靠 kernel 内 event 计时 + 手算等效带宽, 与 reduction 那节同款套路;DRAM Throughput 那行留给 A100(原生 Linux)跑。

6. 为什么这是 AI 算子的提速主线

Transformer 里 softmax 和 LayerNorm 极其密集:每个 attention 有 softmax,每个子层前后有 LayerNorm。 它们都是访存瓶颈的"reduction + 逐元素"算子。单独看每个都已经被框架调到接近带宽极限, 所以现代推理/训练引擎的提速几乎都来自更大范围的融合:

融合省掉的 HBM 往返
FlashAttentionQKᵀ→softmax→·V 全程不落 N×N 的注意力矩阵
fused LayerNorm + 残差residual add 与 norm 合一,中间和不落 HBM
fused bias + GELU逐元素链合并,一趟读写

你这一节手写的"每行一个 block、online 归约、一趟读完"正是 FlashAttention 的核心积木—— 它把整个 attention 拆成分块的 online softmax,从不物化 N×N 矩阵。理解了融合 softmax,你就理解了 FlashAttention 为什么快。

7. 练习:融合决策(混合前几节)

点选项看反馈。这些题把 roofline、带宽、shared、warp shuffle 都串起来了。

💬 随时问我。 想看一份完整可编译的融合 softmax(online + warp shuffle,带 4-kernel 对照基线)? Welford 的 warp 归约具体怎么写?N 超过 shared 容量(比如词表 softmax,N=5万)该怎么分块? FlashAttention 的 online softmax 和这里的有什么不同?把问题抛来。我是你的老师。

主源推荐(本节精读)

📘 Milakov & Gimelshein, "Online normalizer calculation for softmax"(NVIDIA, 2018) ——online softmax 的原始论文,一趟同时算 max 与 sum 的推导就在这,FlashAttention 的软最大化也基于它。