归一化:从特征尺度到残差路径与 CUDA 实现

把 BatchNorm、LayerNorm、RMSNorm 与 Pre-Norm 的数据流串起来

Posted by Liu Mengxuan on September 22, 2026

一个 token 进入 Transformer 后,会被表示成一个有几千维的向量。Attention 和 FFN 不断加工这些特征,残差连接则把每次加工得到的更新加回原表示。随着计算逐层推进,向量的数值尺度也在变化。

归一化为这些计算提供了尺度受控的输入。理解它,可以沿着一条线展开:先确定哪些数放在一起统计,再看这些数怎样被调整,接着把归一化放回残差结构,最后观察 GPU 如何完成同样的计算。

1. 归一化从选定统计范围开始

设 Transformer 的输入形状为 [B, S, D]B 是 batch size,S 是 token 数,D=d_model 是每个 token 的特征维度。归一化不改变这个形状,但不同方法会选择不同的统计范围。

先用一个二维矩阵理解方向:每一行是一个样本,每一列是一个特征。

同一个三行四列数值矩阵,BatchNorm 高亮一列,LayerNorm 高亮一行,表示不同的统计范围

BatchNorm 固定特征维度,跨样本统计。 图中第一列的三个数共同产生一组均值和方差,其他列分别产生自己的统计量。训练时通常使用当前 batch 的统计量,同时积累 running mean 和 variance;推理时通常使用积累的统计量,因此可以处理单样本输入。

Transformer 中的 LayerNorm 固定一个 token,沿它自己的特征维度统计。[B, S, D] 来说,就是沿最后一维 D 计算;每个位置各自产生一个均值和一个方差。

这种局部统计不依赖 batch 中还有哪些样本,也不使用其他位置的特征。句子长短和 padding 不会污染另一个 token 的统计量,训练与推理采用同一套计算规则,因此很适合自回归模型。

2. LayerNorm 调整中心与尺度

取一个 token 的特征向量:

$x=[x_0,x_1,\ldots,x_{D-1}]$

LayerNorm 先计算均值:

$\mu=\frac{1}{D}\sum_{j=0}^{D-1}x_j$

再计算方差,即特征偏离均值的平方的平均:

$v=\frac{1}{D}\sum_{j=0}^{D-1}(x_j-\mu)^2$

最后执行标准化和可学习的仿射变换:

$y_j=\gamma_j\frac{x_j-\mu}{\sqrt{v+\epsilon}}+\beta_j$

减去均值把特征移到以零为中心的位置,除以标准差调整整体尺度。ε 是一个小正数,用于稳定分母。γβ 是长度为 D 的可学习向量,分别调整每个通道的缩放与偏移;同一层的所有 token 共享这两组参数,但各自计算自己的均值和方差。

x=[1,2,3,4] 为例:

$\mu=2.5,\qquad v=1.25$

暂时忽略 ε,并令 γ=1、β=0

1
2
3
原始特征: [ 1,      2,      3,      4    ]
减去均值: [-1.5,   -0.5,    0.5,    1.5  ]
除标准差: [-1.342, -0.447,  0.447,  1.342]

此时均值约为零、方差约为一。这描述的是标准化后的中间结果;再经过训练得到的逐通道 γ、β 后,最终输出不必保持这些统计值。

3. RMSNorm 保留中心,只调整幅度

RMSNorm 同样对每个 token 独立计算,但省去减均值,使用均方根衡量整体幅度:

$\operatorname{RMS}(x)=\sqrt{\frac{1}{D}\sum_{j=0}^{D-1}x_j^2}$

常见的 RMSNorm 只保留可学习缩放参数:

$y_j=\gamma_j\frac{x_j}{\sqrt{\frac{1}{D}\sum_k x_k^2+\epsilon}}$

对同一个 [1,2,3,4],平方均值为 7.5,均方根约为 2.739。忽略 ε,令 γ=1,得到 [0.365,0.730,1.095,1.461]

数轴比较原始向量、LayerNorm 与 RMSNorm 的四个特征值:LayerNorm 移到零点两侧,RMSNorm 保留正值比例并缩小幅度

图中每一种颜色始终代表同一个特征。LayerNorm 改变中心后,前两个特征变成负值;RMSNorm 在乘 γ 之前只是除以同一个正数,保留比例和符号,改变整体幅度。忽略 ε 时,此阶段的 RMS 为一,均值不必为零。

两种统计量满足:

$\operatorname{RMS}(x)^2=\operatorname{Var}(x)+\operatorname{mean}(x)^2$

均值接近零时,RMS 接近标准差,两种归一化也较接近。但 RMSNorm 不要求输入均值为零;它的设计重点是控制幅度,并以更简单的统计运算支持模型训练。LLaMA 等模型采用了这种结构。

4. Norm 的位置决定残差保留什么

把 Attention 或 FFN 统一记作 F。Post-Norm 与 Pre-Norm 分别是:

$\text{Post-Norm:}\quad y=\operatorname{Norm}(x+F(x))$

$\text{Pre-Norm:}\quad y=x+F(\operatorname{Norm}(x))$

Post-Norm 先加工、残差相加,再归一化整个结果。Pre-Norm 则只在加工分支入口归一化,原始输入沿残差分支直接参与相加。

把两个连续的 Pre-Norm 子层写出来,区别会更清楚:

1
2
3
4
5
6
# Attention 子层
h = x + attention(norm1(x))

# FFN 子层
z = norm2(h)
y = h + ffn(z)

这里 h 相加之后,进入 FFN 前确实又经历了 Norm,但 norm2(h) 只是产生了加工分支所需的 z,并没有替换原来的 h。最后相加使用的仍然是原始 h

用实际数值向量展示 Pre-Norm 保留原始残差 h,同时生成归一化向量 z;Post-Norm 则把归一化后的向量作为后续残差表示

图中橙色向量代表未经归一化的残差表示,蓝色向量代表归一化后的表示。沿着残差分支追踪数值,才能准确判断 Pre/Post;只看计算分支上加法与 Norm 的先后顺序,容易混淆。

Post-Norm 的连续子层则是:

1
2
h = norm1(x + attention(x))
y = norm2(h + ffn(h))

第一次相加的结果必须穿过 norm1,才成为后续传递的 h。下一次残差使用的是这个已经归一化的表示。

5. 残差路径同时也是梯度路径

对 Pre-Norm,令 G(x)=F(Norm(x)),就有:

$y=x+G(x),\qquad \frac{\partial y}{\partial x}=I+J_G$

I 是单位矩阵,来自直接相加的 xJ_G 是加工分支的导数矩阵。反向传播时,除了经过加工分支,还存在一条绕过该分支 Norm 的直接路径。

Post-Norm 则满足:

$\frac{\partial y}{\partial x}=J_{\mathrm{Norm}}(I+J_F)$

两条分支汇合后都经过 Norm,梯度返回时也必须经过它。深层堆叠时,这种结构差异使 Pre-Norm 通常更容易训练稳定,不过并不保证总梯度始终不变,也不代表可以普遍省略 warmup。

Pre-Norm 的残差表示本身仍可能随深度增长。它控制的是每次送进 Attention 或 FFN 的输入尺度;常见架构还会在所有 Block 之后放置最终的 Norm,再进行词表投影。

6. 一行特征在 CUDA 中的分工

在 GPU 上,可以把 [B,S,D] 看成 B×S 行,每行对应一个 token。常见 kernel 让一个线程块处理一行,块内线程分担这行特征;实际分工会根据 D、行数和硬件调整。

例如 D=4096,一个线程块有 256 个线程,每个线程平均处理 16 个元素:

1
2
3
4
线程 0:下标 0、256、512……
线程 1:下标 1、257、513……
……
线程 255:下标 255、511、767……

同一轮中,相邻线程读取相邻元素,有利于合并内存访问。每个线程先做局部统计,再把部分结果归约成整行统计量。

以十六个特征和四个线程示意 RMSNorm:按下标交错分工,各线程计算局部平方和,再合并得到整行缩放系数

图中用 4 个线程处理 16 个特征来展示同样的分工。对于 RMSNorm,每个线程计算局部平方和,再合并为:

$S=\sum_jx_j^2,\qquad r=\operatorname{rsqrt}(S/D+\epsilon)$

rsqrt 表示平方根的倒数。得到整行共用的 r 后,各线程独立计算自己的输出:

$y_j=\gamma_jx_jr$

归约通常先在 warp 内完成,再按需要通过共享内存合并不同 warp 的结果。LayerNorm 还需要均值和方差,可以采用并行 Welford:每个线程维护元素数量、均值和平方偏差总和,再按合并公式组合统计量。

Welford 避免直接使用 mean(x²)−mean(x)² 时的大数相减。例如 [1000,1001] 的方差是 0.25,却可能要由两个百万量级的数相减得到,有限精度下很容易损失有效数字。

7. 将残差融合与混合精度落到同一份数据上

统计结束后,每个输出仍需要原始特征。如果线程把输入保存在寄存器里,就可以直接复用,减少再次加载。代价是寄存器占用增加,需要与并行运行能力权衡。RMSNorm 的统计更简单,但两种 Norm 都可以利用片上缓存,因此不能仅凭公式断言固定少一次 HBM 读取。

这一思路还能应用于相邻的残差加法与归一化:

1
2
3
h = residual + sublayer_output
z = norm(h)
y = h + next_sublayer(z)

融合 kernel 读取两个加数后,在寄存器中计算 h,立即用它做归约和归一化,最后写出 hz。这里有两个不同用途的输出:h 保留给残差分支,z 送入加工分支。因而主要省去的是 Norm 对 h 的重新读取,以及一次 kernel 启动;h 的写回仍然需要保留。

数据存储通常采用 FP16/BF16,而平方、统计和归一化中的关键计算常采用 FP32。顺序尤其重要:先提升精度,再平方和累加。FP16 可以保存 300,却不能保存其平方 90000;若先平方溢出,再转 FP32,已经丢失的信息无法恢复。

1
2
3
4
5
6
# RMSNorm 的教学实现;实际高效 kernel 会融合这些操作
x32 = x.float()
mean_square = x32.square().mean(dim=-1, keepdim=True)
inv_rms = torch.rsqrt(mean_square + eps)
y32 = x32 * inv_rms * weight.float()
y = y32.to(x.dtype)

mean_squareinv_rms 的形状都是 [B,S,1],通过广播作用于当前 token 的所有特征。BF16 虽然具有更大的动态范围,仍需要关注累加误差,因此也常使用 FP32 归约。

从公式到实现,始终处理的是同一行特征:先明确统计范围,再产生共享统计量,最后逐元素输出。Pre-Norm 的结构则决定哪些数据必须保留——归一化后的计算输入和未经归一化的残差表示,各自承担不同的作用。

参考资料