CUDA 算子工程:手写 FlashAttention v2 之路
第 7 章 LayerNorm 与 RMSNorm
方差有两个数学上等价、数值上天差地别的算法。 这一章讲清楚为什么 在 FP32 上会失效,而 Welford 的增量公式不会。
7.1 为什么要归一化
Transformer 里每个 block 都套着一层归一化:Pre-LN 或 Post-LN。它的作用是把每个 token 的隐藏维度(hidden_size,4096 / 8192 / ...)的激活值"拉到合适的尺度",防止训练时数值爆炸或塌陷。
LayerNorm 的标准定义:
其中:
- 是均值
- 是方差
- 是可学习的缩放与偏置(每个特征维独立)
- 是数值稳定的小量(典型 )
每次 attention 之前一次 LayerNorm,每次 FFN 之前一次 LayerNorm——一个 32 层的 Transformer,每次 forward 要做 64+ 次 LayerNorm。这是 LLM 推理中除 GEMM 和 attention 之外最频繁的算子之一。
7.2 朴素两遍算法
最直观的写法:
// Pass 1: 算 mean
float mean = 0;
for (int i = 0; i < H; ++i) mean += x[i];
mean /= H;
// Pass 2: 算 variance
float var = 0;
for (int i = 0; i < H; ++i) var += (x[i] - mean) * (x[i] - mean);
var /= H;
// Pass 3: 归一化
float rstd = rsqrtf(var + eps);
for (int i = 0; i < H; ++i)
y[i] = (x[i] - mean) * rstd * gamma[i] + beta[i];
3 pass,每 pass 都遍历 x 一次。在 GPU 上意味着 x 数据从 HBM 读 3 次——和 safe softmax 一样的痛点。
7.3 一个数学上的等价:方差的两种公式
数学上有这样一个等式:
这意味着方差可以用"和"与"平方和"两个量计算:
如果用这个公式,可以用一遍同时收集 和 :
// 1 pass: 算 sum 和 sum_sq
float sum = 0, sum_sq = 0;
for (int i = 0; i < H; ++i) {
sum += x[i];
sum_sq += x[i] * x[i];
}
float mean = sum / H;
float var = sum_sq / H - mean * mean;
// Pass 2: 归一化 (这一遍不可避免, 因为输出依赖 mean/var)
float rstd = rsqrtf(var + eps);
for (int i = 0; i < H; ++i)
y[i] = (x[i] - mean) * rstd * gamma[i] + beta[i];
成功——从 3 pass 降到 2 pass!但有一个致命的数值稳定问题:
这个公式叫naive 方差公式,它在数值上非常不稳定。当 的均值很大、方差很小时, 和 会接近相等,相减时会发生灾难性消除(catastrophic cancellation),损失大量有效数字。
举个具体例子:
x = [1000.001, 1000.002, 1000.003, ..., 1000.010] (10 个数, 等差 0.001)
真实 mean = 1000.0055
真实 var = 8.25e-6 (偏差平方和 8.25e-5, 除以 10)
朴素公式:
sum = 10000.055
sum_sq = 10000110.000385
mean = 1000.0055, mean^2 = 1000011.00003025
var = (10000110.000385 / 10) - mean^2
= 1000011.0000385 - 1000011.0000302...
= 8.25e-6
——两个要相减的数都在 1e6 量级, 而答案在 1e-6 量级:
差了 12 个数量级。FP32 只有约 7 位十进制有效数字,
1e6 附近相邻两个 FP32 之间就差 0.0625, 那个 8.25e-6
早就落在末位以下了 —— 按 FP32 实测, 随累加顺序不同
算出来是 0 或 0.25 (真值的约 3 万倍)。
这就是为什么真实场景下不能用这个公式——输入的 magnitude 一旦大了,FP32 精度根本不够。
LLM 里激活值出现上百量级的离群值并不罕见;一旦某一行的均值相对标准差很大,用朴素方差公式做 LayerNorm,方差就可能算成负数(开方变 NaN)或者完全错误的值。
7.4 Welford 算法:数值稳定的一遍方差
1962 年统计学家 B. P. Welford 提出了一个数值稳定的在线方差算法。它的核心是维护"当前均值"和"M2(偏差平方和)"两个状态,用增量更新:
定义:
- = 已处理元素数量
- = 前 个元素的均值
- = 偏差平方和
初始 , , 。每来一个新元素 :
最后方差 。
Welford 在数值上是稳定的——它不像朴素公式那样减两个相近的大数,而是不断累积小的偏差量。FP32 下也能保持精度。
7.4.1 Welford 的并行合并规则
Welford 的妙处是它有一个可结合的合并规则。如果两个独立计算的 partial Welford 状态 和 要合并:
这个合并是对称、可结合的。和 online softmax 一样,它构成了一个 monoid——可以用并行 reduce 来计算。
7.4.2 Welford 的 GPU kernel 形式
struct WelfordState {
int n;
float mean;
float m2;
};
__device__ WelfordState welford_update(WelfordState s, float x) {
s.n += 1;
float delta = x - s.mean;
s.mean += delta / s.n;
s.m2 += delta * (x - s.mean);
return s;
}
__device__ WelfordState welford_combine(WelfordState a, WelfordState b) {
int n = a.n + b.n;
if (n == 0) return {0, 0, 0};
float delta = b.mean - a.mean;
float new_mean = a.mean + delta * b.n / n;
float new_m2 = a.m2 + b.m2 + delta * delta * a.n * b.n / n;
return {n, new_mean, new_m2};
}
把这两个函数用 warp shuffle 串起来:
__device__ WelfordState warp_welford_reduce(WelfordState s) {
for (int offset = 16; offset > 0; offset >>= 1) {
WelfordState other;
other.n = __shfl_xor_sync(0xFFFFFFFF, s.n, offset);
other.mean = __shfl_xor_sync(0xFFFFFFFF, s.mean, offset);
other.m2 = __shfl_xor_sync(0xFFFFFFFF, s.m2, offset);
s = welford_combine(s, other);
}
return s;
}
这就是 Apex FusedLayerNorm 内部的核心。第 7.6 节会给完整的 kernel。
7.5 RMSNorm:进一步简化
2019 年 Zhang & Sennrich 在论文 Root Mean Square Layer Normalization 中提出了 RMSNorm——一个比 LayerNorm 更简单的归一化方案。LLaMA、Mistral、Qwen、Gemma 等几乎所有现代开源 LLM 都用 RMSNorm。
RMSNorm 的定义:
对比 LayerNorm,RMSNorm 去掉了:
- 均值减法(不再需要 )
- 均值统计量(不再需要 )
- bias (只保留 )
为什么这样改?论文给出的理由是"重新中心化(减均值)对模型表现影响很小,但会增加计算"。RMSNorm 论文自己在多种网络和任务上测得的是质量与 LayerNorm 相当、运行时间减少 7%~64%——上限取决于 Norm 在该模型里占多大计算比重。LLaMA 之后的大模型普遍采用它,验证了大规模下质量同样不掉。
7.5.1 RMSNorm 的算法只需要 1 个状态量
LayerNorm 的 Welford 需要维护 (n, mean, m2) 三个状态。RMSNorm 只需要维护平方和:
float ss = 0;
for (int i = tid; i < H; i += blockDim.x) {
float v = x[i];
ss += v * v;
}
// ss 在 block 内跨线程归约 (同 7.6 节的 warp + SMEM 两级)
float inv_rms = rsqrtf(ss / H + eps);
for (int i = tid; i < H; i += blockDim.x) {
y[i] = x[i] * inv_rms * gamma[i];
}
这个 kernel 比 LayerNorm 简单得多——纯标量 reduce,无需 Welford 那套增量更新。
7.5.2 性能对比
LayerNorm 与 RMSNorm 的每 token 开销对比(量级估算):
| 算子 | 每 token 浮点操作数 | 每 token 状态量 | HBM 流量 |
|---|---|---|---|
| LayerNorm (Welford) | ~6H | (n, mean, m2) 三个 |
读 1 遍 + 写 1 遍 |
| RMSNorm | ~3H | 平方和一个 | 读 1 遍 + 写 1 遍 |
两者的 HBM 流量是一样的(都是带宽 bound、都读一遍写一遍),差别出在算术与寄存器开销:Welford 的每元素更新里有一次除法,而 RMSNorm 只有一次 FMA。所以在 H 不大、kernel 还没完全被带宽压住的时候差距明显,H 很大时两者会一起顶到 HBM 峰值附近。
7.6 完整的 Fused LayerNorm Kernel
把上面的元素拼起来,给一个完整的、以逼近 HBM 峰值为目标的 LayerNorm kernel(假设 H 是 4 的倍数且各指针 16 字节对齐,VEC_SIZE 只能取 4,因为读写写死了 float4):
template <int BLOCK_SIZE = 512, int VEC_SIZE = 4>
__global__ void layernorm_fwd(
const float* __restrict__ x, // [B, H]
const float* __restrict__ gamma, // [H]
const float* __restrict__ beta, // [H]
float* __restrict__ y, // [B, H]
float* __restrict__ mean_out, // [B]
float* __restrict__ rstd_out, // [B]
int H,
float eps
) {
int row = blockIdx.x;
int tid = threadIdx.x;
const float* x_row = x + row * H;
float* y_row = y + row * H;
// ============ Phase 1: Welford reduce ============
WelfordState state = {0, 0.0f, 0.0f};
// 每线程读 VEC_SIZE 个元素。前提: H % 4 == 0 且指针 16B 对齐,
// 否则要补标量尾部分支 (PyTorch 此时直接退回非向量化路径)
for (int i = tid * VEC_SIZE; i < H; i += BLOCK_SIZE * VEC_SIZE) {
float4 v = *reinterpret_cast<const float4*>(&x_row[i]);
state = welford_update(state, v.x);
state = welford_update(state, v.y);
state = welford_update(state, v.z);
state = welford_update(state, v.w);
}
// Warp 内合并
state = warp_welford_reduce(state);
// Block 内合并 (跨 warp)
__shared__ WelfordState warp_states[BLOCK_SIZE / 32];
int warp_id = tid / 32;
int lane_id = tid % 32;
if (lane_id == 0) warp_states[warp_id] = state;
__syncthreads();
if (warp_id == 0) {
if (lane_id < BLOCK_SIZE / 32) state = warp_states[lane_id];
else state = {0, 0.0f, 0.0f};
for (int offset = 16; offset > 0; offset >>= 1) {
WelfordState other;
other.n = __shfl_xor_sync(0xFFFFFFFF, state.n, offset);
other.mean = __shfl_xor_sync(0xFFFFFFFF, state.mean, offset);
other.m2 = __shfl_xor_sync(0xFFFFFFFF, state.m2, offset);
state = welford_combine(state, other);
}
}
__shared__ float final_mean, final_rstd;
if (warp_id == 0 && lane_id == 0) {
final_mean = state.mean;
float var = state.m2 / state.n;
final_rstd = rsqrtf(var + eps);
if (mean_out) mean_out[row] = final_mean;
if (rstd_out) rstd_out[row] = final_rstd;
}
__syncthreads();
// ============ Phase 2: Normalize + scale ============
for (int i = tid * VEC_SIZE; i < H; i += BLOCK_SIZE * VEC_SIZE) {
float4 xv = *reinterpret_cast<const float4*>(&x_row[i]);
float4 gv = *reinterpret_cast<const float4*>(&gamma[i]);
float4 bv = *reinterpret_cast<const float4*>(&beta[i]);
float4 yv;
yv.x = (xv.x - final_mean) * final_rstd * gv.x + bv.x;
yv.y = (xv.y - final_mean) * final_rstd * gv.y + bv.y;
yv.z = (xv.z - final_mean) * final_rstd * gv.z + bv.z;
yv.w = (xv.w - final_mean) * final_rstd * gv.w + bv.w;
*reinterpret_cast<float4*>(&y_row[i]) = yv;
}
}
关键点:
- 每个 block 处理一行:grid_size = batch_size。
- VEC_SIZE=4 用 float4 vectorized I/O:减少访存指令条数(nvcc 13.4 编译为
LDG.E.128/STG.E.128,不同 CUDA 版本可能不同)。 - Welford state 在 warp 内、block 内归约:用 shfl 和 SMEM。
- 存 mean/rstd 给反向用:反向 LN 需要这两个量。
这个结构(一行一个 block、float4 读写、Welford 在 warp/block 内两级归约)在 HBM 层面接近只读一遍写一遍:Phase 2 对 x 的第二次读在代码上仍是一次全局访存,但这一行刚读过,通常命中 L1/L2——PyTorch 的 vectorized_layer_norm_kernel 也是这样先 compute_stats 再重读 X。HBM 流量因此接近 LayerNorm 应有的下界,剩下的差距在指令调度和 occupancy 上,目标是逼近 HBM 峰值带宽。
本专栏没有条件在写作时对上面这个 kernel 与 PyTorch / Apex 做同机对比,所以这里不给具体的 GB/s 数字。要比就照本章末的练习自己在目标卡上跑——不同的 B/H 组合、不同的 dtype,排名是会变的。
7.7 反向:LayerNorm 的两个公式
LayerNorm 反向比前向复杂得多。给定 ,需要算:
其中 ,,,。注意 只乘在 上(它是逐维的),不能提到整个括号外面——这是手写反向时最常见的一个错。
这个反向公式有两个 reduce: 和 。可以一遍同时算两个 reduce——把两个和打包成一对一起做 warp/block 归约;它们只是普通求和,不需要 Welford 那套增量更新。
完整反向 kernel 比较长,本专栏不展开。读者可以参考 Apex 的 layer_norm_cuda_kernel.cu,那是工业级的参考实现。
RMSNorm 反向比 LayerNorm 简单得多——对 的梯度只需要一个行内 reduce:。这是 RMSNorm 在训练中的另一个加速点。
7.8 与 PyTorch / Apex 实现对照
PyTorch 的 LayerNorm 实现在 pytorch-v2.11.0/aten/src/ATen/native/cuda/layer_norm_kernel.cu,v2.11 里有两条路径,分派条件写在 :1111:dtype 是 float / at::Half / at::BFloat16、N 不超过 (源码里写的是 1ULL << std::numeric_limits<float>::digits,因为 count 用 float 存)、N 是向量宽度的整数倍、且 X / Y / gamma / beta 四个指针都按 vec_size * sizeof(T) 对齐——全满足才走 vectorized_layer_norm_kernel(:343,一个 kernel 里 compute_stats + 归一化一起做完);任何一条不满足就退回 RowwiseMomentsCUDAKernel(:59)+ LayerNormForwardCUDAKernel(:101)两个 kernel。两条路都是 per-row + Welford + shuffle 合并,结构和上面的 layernorm_fwd 一致:快路径用本文件自带的 WelfordDataLN(:126)/ cuWelfordOnlineSum(:135)/ cuWelfordCombine(:155);回退路径的 RowwiseMomentsCUDAKernel 则用通用的 WelfordOps(pytorch-v2.11.0/aten/src/ATen/native/SharedReduceOps.h,合并规则在 combine,:118)配 cuda_utils::BlockReduce。值得一提的是 PyTorch 和 Apex 在这里走了同一条路子:整个文件用一个 rms_norm 模板参数把 RMSNorm 复用进同一套 kernel。
Apex 的 FusedLayerNorm 在 apex-25.09/csrc/layer_norm_cuda_kernel.cu,做了几个额外优化:
- mixed precision:
mu/sigma2/count用at::acc_type<scalar_t_in, true>累加(cuWelfordMuSigma2,:70;类型在 :990 选定),输入是 fp16 / bf16 / fp32 时一律是float(double 输入则是double),避免半精度累加掉精度。 - 只对
at::Half做了专门特化(:179):先检查((size_t)lvals) & 3,未对齐时让 0 号线程单独消化掉首元素凑齐 32 位对齐,之后所有线程按__half2一次吃 8 个元素(for (; l+7 < n2; l += 8*numx),:221)。注意at::BFloat16没有这个特化,走的是通用模板的 4 元素展开(:97)。 - 同一份 kernel 用
rms_only开关同时支持 LayerNorm 和 RMSNorm:rms_only为真时走cuRMSOnlineSum(:53),整套均值统计被跳过。
warp 内和跨 warp 的合并用的都是 Chan 的并行公式(源码里叫 cuChanOnlineSum,:28),和 7.4.1 节给的合并规则是同一个东西。至于"Apex 和 PyTorch 谁快"——这要看 dtype、H、batch 和卡,本专栏不给一个笼统的百分比;绝大多数情况下 PyTorch 默认实现已经够用。
7.9 工程权衡:Pre-LN vs Post-LN
最后顺便说一下 LayerNorm 的两种位置安排,因为它影响 kernel 调度:
flowchart LR
subgraph PostLN [Post-LN · 原始 Transformer]
P1[x] --> P2[Attention]
P2 --> P3[+ x]
P3 --> P4[LayerNorm]
P4 --> P5[FFN]
P5 --> P6[+ ...]
P6 --> P7[LayerNorm]
end
subgraph PreLN [Pre-LN · 现代 LLM]
L1[x] --> L2[LayerNorm]
L2 --> L3[Attention]
L3 --> L4[+ x]
L4 --> L5[LayerNorm]
L5 --> L6[FFN]
L6 --> L7[+ ...]
end
Post-LN(原始 2017 Transformer):LayerNorm 在残差之后。训练困难(梯度消失/爆炸),现在很少用。
Pre-LN(GPT-2/LLaMA/几乎所有现代 LLM):LayerNorm 在残差之前。训练稳定,但有"残差累积"问题(不严重)。
工程上 Pre-LN 还有一个kernel fusion 优势:Pre-LN 的输出只被接下来的 GEMM 消费(残差流走的是另一条支路),所以Norm + 接下来的 GEMM 输入可以融合,省掉一次以高精度写回 HBM 再读回。实际工程里常见的是它的退一步形态:把 Norm 与 GEMM 输入的量化合成一个 kernel,Norm 直接输出 FP8 供 GEMM 读——vLLM v0.8.5 的 vllm-0.8.5/vllm/compilation/fusion.py 就把 rms_norm / fused_add_rms_norm 后接 FP8 量化的模式替换成 rms_norm_static_fp8_quant 等融合算子(实现在 vllm-0.8.5/csrc/layernorm_quant_kernels.cu)。
7.10 这一章的小结与下一章
这一章的关键收获:
- 方差有两个公式:朴素 数值不稳定;Welford 增量公式数值稳定。
- Welford 算法和 online softmax 同源——都是用一个可结合 monoid 的 combine 规则把多 pass 改成 1 pass。
- RMSNorm 是 LayerNorm 的简化:去掉均值统计,只算 RMS,状态量从三个降到一个。论文报告的运行时间下降是 7%~64%(取决于 Norm 在模型里的占比),效果几乎无损。这就是为什么主流开源 LLM 普遍改用 RMSNorm。
- Fused LayerNorm 的目标是逼近 HBM 峰值带宽——这是带宽 bound kernel 的合理目标。
- Pre-LN + GEMM Fusion 是工业级推理引擎的常见优化。
第 8 章我们继续往工程化深入——讲 Element-wise 算子的融合。LLM 推理里有大量"小算子"(add、mul、ReLU、GeLU、SiLU、dropout),单独跑每个都是带宽 bound、利用率极低。把它们融合到一起跑(或者融合到 LayerNorm/GEMM 里),是减少 HBM 往返的关键手段。读完第 8 章读者会理解 vLLM 的 fused_add_rms_norm kernel 到底省掉了哪几次 HBM 往返和哪一次 kernel launch。
本章动手练习:
- 实现两版 LayerNorm:朴素方差公式版本和 Welford 版本,输入用 7.3 节
[1000.001, 1000.002, ...]这种大均值小方差的数据(别再把步长缩到 0.0001:1000 附近 FP32 的间距约 6.1e-5,输入本身就表示不准),对比两者的精度。- 写一个 RMSNorm kernel,对比 LayerNorm 在 H=8192 上的实测延迟。
- 阅读 PyTorch 的
RowwiseMomentsCUDAKernel(pytorch-v2.11.0/aten/src/ATen/native/cuda/layer_norm_kernel.cu:59),找到 Welford combine 规则在源码里的对应行(提示:它用的是pytorch-v2.11.0/aten/src/ATen/native/SharedReduceOps.h里WelfordOps::combine,:118;同文件快路径用的cuWelfordCombine在 :155)。