CUDA 算子工程:手写 FlashAttention v2 之路

第 7 章 LayerNorm 与 RMSNorm

作者 杨艺韬 · 4,032 字 · 发布于 · 更新于

方差有两个数学上等价、数值上天差地别的算法。 这一章讲清楚为什么 E[X2]−(E[X])2E[X^2]-(E[X])^2 在 FP32 上会失效,而 Welford 的增量公式不会。

7.1 为什么要归一化

Transformer 里每个 block 都套着一层归一化:Pre-LN 或 Post-LN。它的作用是把每个 token 的隐藏维度(hidden_size,4096 / 8192 / ...)的激活值"拉到合适的尺度",防止训练时数值爆炸或塌陷。

LayerNorm 的标准定义:

LN(x)=x−μσ2+ϵ⋅γ+β\text{LN}(\mathbf{x}) = \frac{\mathbf{x} - \mu}{\sqrt{\sigma^2 + \epsilon}} \cdot \gamma + \beta

其中:

  • μ=1H∑i=1Hxi\mu = \frac{1}{H} \sum_{i=1}^{H} x_i 是均值
  • σ2=1H∑i=1H(xi−μ)2\sigma^2 = \frac{1}{H} \sum_{i=1}^{H} (x_i - \mu)^2 是方差
  • γ,β\gamma, \beta 是可学习的缩放与偏置(每个特征维独立)
  • ϵ\epsilon 是数值稳定的小量(典型 10−510^{-5})

每次 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 一个数学上的等价:方差的两种公式

数学上有这样一个等式:

σ2=E[X2]−(E[X])2\sigma^2 = E[X^2] - (E[X])^2

这意味着方差可以用"和"与"平方和"两个量计算:

σ2=1H∑ixi2−μ2\sigma^2 = \frac{1}{H} \sum_{i} x_i^2 - \mu^2

如果用这个公式,可以用一遍同时收集 ∑xi\sum x_i 和 ∑xi2\sum x_i^2:

// 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!但有一个致命的数值稳定问题:

E[X2]−(E[X])2E[X^2] - (E[X])^2 这个公式叫naive 方差公式,它在数值上非常不稳定。当 XX 的均值很大、方差很小时,E[X2]E[X^2] 和 (E[X])2(E[X])^2 会接近相等,相减时会发生灾难性消除(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(偏差平方和)"两个状态,用增量更新:

定义:

  • nn = 已处理元素数量
  • μn\mu_n = 前 nn 个元素的均值
  • M2,n=∑i=1n(xi−μn)2M_{2,n} = \sum_{i=1}^{n} (x_i - \mu_n)^2 = 偏差平方和

初始 n=0n=0, μ=0\mu = 0, M2=0M_2 = 0。每来一个新元素 xn+1x_{n+1}:

n′=n+1δ=xn+1−μnμn′=μn+δ/n′M2,n′=M2,n+δ⋅(xn+1−μn′)\begin{aligned} n' &= n + 1 \\ \delta &= x_{n+1} - \mu_n \\ \mu_{n'} &= \mu_n + \delta / n' \\ M_{2, n'} &= M_{2, n} + \delta \cdot (x_{n+1} - \mu_{n'}) \end{aligned}

最后方差 σ2=M2,N/N\sigma^2 = M_{2, N} / N。

Welford 在数值上是稳定的——它不像朴素公式那样减两个相近的大数,而是不断累积小的偏差量。FP32 下也能保持精度。

7.4.1 Welford 的并行合并规则

Welford 的妙处是它有一个可结合的合并规则。如果两个独立计算的 partial Welford 状态 (na,μa,M2,a)(n_a, \mu_a, M_{2,a}) 和 (nb,μb,M2,b)(n_b, \mu_b, M_{2,b}) 要合并:

n=na+nbδ=μb−μaμ=μa+δ⋅nb/nM2=M2,a+M2,b+δ2⋅na⋅nb/n\begin{aligned} n &= n_a + n_b \\ \delta &= \mu_b - \mu_a \\ \mu &= \mu_a + \delta \cdot n_b / n \\ M_2 &= M_{2,a} + M_{2,b} + \delta^2 \cdot n_a \cdot n_b / n \end{aligned}

这个合并是对称、可结合的。和 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 的定义:

RMS(x)=x1H∑ixi2+ϵ⋅γ\text{RMS}(\mathbf{x}) = \frac{\mathbf{x}}{\sqrt{\frac{1}{H}\sum_i x_i^2 + \epsilon}} \cdot \gamma

对比 LayerNorm,RMSNorm 去掉了:

  1. 均值减法(不再需要 x−μ\mathbf{x} - \mu)
  2. 均值统计量(不再需要 μ\mu)
  3. bias β\beta(只保留 γ\gamma)

为什么这样改?论文给出的理由是"重新中心化(减均值)对模型表现影响很小,但会增加计算"。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;
    }
}

关键点:

  1. 每个 block 处理一行:grid_size = batch_size。
  2. VEC_SIZE=4 用 float4 vectorized I/O:减少访存指令条数(nvcc 13.4 编译为 LDG.E.128 / STG.E.128,不同 CUDA 版本可能不同)。
  3. Welford state 在 warp 内、block 内归约:用 shfl 和 SMEM。
  4. 存 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 反向比前向复杂得多。给定 dy=∂L/∂ydy = \partial L / \partial y,需要算:

∂L∂γi=∑ndyi(n)⋅x^i(n)∂L∂βi=∑ndyi(n)∂L∂xi=1σ2+ϵ(gi−gˉ−x^i⋅gx^‾)\begin{aligned} \frac{\partial L}{\partial \gamma_i} &= \sum_n dy^{(n)}_i \cdot \hat{x}^{(n)}_i \\ \frac{\partial L}{\partial \beta_i} &= \sum_n dy^{(n)}_i \\ \frac{\partial L}{\partial x_i} &= \frac{1}{\sqrt{\sigma^2 + \epsilon}} \left( g_i - \bar{g} - \hat{x}_i \cdot \overline{g \hat{x}} \right) \end{aligned}

其中 x^i=(xi−μ)/σ2+ϵ\hat{x}_i = (x_i - \mu) / \sqrt{\sigma^2 + \epsilon},gi=dyiγig_i = dy_i \gamma_i,gˉ=1H∑jgj\bar{g} = \frac{1}{H}\sum_j g_j,gx^‾=1H∑jgjx^j\overline{g\hat{x}} = \frac{1}{H}\sum_j g_j \hat{x}_j。注意 γi\gamma_i 只乘在 dyidy_i 上(它是逐维的),不能提到整个括号外面——这是手写反向时最常见的一个错。

这个反向公式有两个 reduce:∑gi\sum g_i 和 ∑gix^i\sum g_i \hat{x}_i。可以一遍同时算两个 reduce——把两个和打包成一对一起做 warp/block 归约;它们只是普通求和,不需要 Welford 那套增量更新。

完整反向 kernel 比较长,本专栏不展开。读者可以参考 Apex 的 layer_norm_cuda_kernel.cu,那是工业级的参考实现。

RMSNorm 反向比 LayerNorm 简单得多——对 xx 的梯度只需要一个行内 reduce:∑jdyjγjxj\sum_j dy_j \gamma_j x_j。这是 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 不超过 2242^{24}(源码里写的是 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,做了几个额外优化:

  1. mixed precision:mu / sigma2 / count 用 at::acc_type<scalar_t_in, true> 累加(cuWelfordMuSigma2,:70;类型在 :990 选定),输入是 fp16 / bf16 / fp32 时一律是 float(double 输入则是 double),避免半精度累加掉精度。
  2. 只对 at::Half 做了专门特化(:179):先检查 ((size_t)lvals) & 3,未对齐时让 0 号线程单独消化掉首元素凑齐 32 位对齐,之后所有线程按 __half2 一次吃 8 个元素(for (; l+7 < n2; l += 8*numx),:221)。注意 at::BFloat16 没有这个特化,走的是通用模板的 4 元素展开(:97)。
  3. 同一份 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 这一章的小结与下一章

这一章的关键收获:

  1. 方差有两个公式:朴素 E[X2]−(E[X])2E[X^2] - (E[X])^2 数值不稳定;Welford 增量公式数值稳定。
  2. Welford 算法和 online softmax 同源——都是用一个可结合 monoid 的 combine 规则把多 pass 改成 1 pass。
  3. RMSNorm 是 LayerNorm 的简化:去掉均值统计,只算 RMS,状态量从三个降到一个。论文报告的运行时间下降是 7%~64%(取决于 Norm 在模型里的占比),效果几乎无损。这就是为什么主流开源 LLM 普遍改用 RMSNorm。
  4. Fused LayerNorm 的目标是逼近 HBM 峰值带宽——这是带宽 bound kernel 的合理目标。
  5. 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。

本章动手练习:

  1. 实现两版 LayerNorm:朴素方差公式版本和 Welford 版本,输入用 7.3 节 [1000.001, 1000.002, ...] 这种大均值小方差的数据(别再把步长缩到 0.0001:1000 附近 FP32 的间距约 6.1e-5,输入本身就表示不准),对比两者的精度。
  2. 写一个 RMSNorm kernel,对比 LayerNorm 在 H=8192 上的实测延迟。
  3. 阅读 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)。