一个 token 进入 Transformer 后,会被表示成一个有几千维的向量。Attention 和 FFN 不断加工这些特征,残差连接则把每次加工得到的更新加回原表示。随着计算逐层推进,向量的数值尺度也在变化。
归一化为这些计算提供了尺度受控的输入。理解它,可以沿着一条线展开:先确定哪些数放在一起统计,再看这些数怎样被调整,接着把归一化放回残差结构,最后观察 GPU 如何完成同样的计算。
1. 归一化从选定统计范围开始
设 Transformer 的输入形状为 [B, S, D]:B 是 batch size,S 是 token 数,D=d_model 是每个 token 的特征维度。归一化不改变这个形状,但不同方法会选择不同的统计范围。
先用一个二维矩阵理解方向:每一行是一个样本,每一列是一个特征。
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 在乘 γ 之前只是除以同一个正数,保留比例和符号,改变整体幅度。忽略 ε 时,此阶段的 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/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 是单位矩阵,来自直接相加的 x;J_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……
同一轮中,相邻线程读取相邻元素,有利于合并内存访问。每个线程先做局部统计,再把部分结果归约成整行统计量。
图中用 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,立即用它做归约和归一化,最后写出 h 与 z。这里有两个不同用途的输出: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_square 与 inv_rms 的形状都是 [B,S,1],通过广播作用于当前 token 的所有特征。BF16 虽然具有更大的动态范围,仍需要关注累加误差,因此也常使用 FP32 归约。
从公式到实现,始终处理的是同一行特征:先明确统计范围,再产生共享统计量,最后逐元素输出。Pre-Norm 的结构则决定哪些数据必须保留——归一化后的计算输入和未经归一化的残差表示,各自承担不同的作用。
参考资料
- Batch Normalization
- Layer Normalization
- Root Mean Square Layer Normalization
- On Layer Normalization in the Transformer Architecture
- PyTorch LayerNorm