从零手写大模型 · 推理篇 08 - LayerNorm(层归一化)
推理篇 · 查看系列继续用"猫吃鱼"这个例子,拆解 LayerNorm 公式背后的几何直觉。
一、LayerNorm 在 Pre-LN 结构里的位置
Transformer 的残差子层,我们用的是 Pre-LN 结构(和 GPT-2 保持一致):
x1 = x + Attention(LayerNorm(x))
LayerNorm 在这里的作用,是在数据进入 Attention 子层之前,先对它做一次"清洗"。这一节要讲清楚的问题是:这个"清洗"具体清洗的是什么?
二、公式回顾:减均值,除标准差
LayerNorm 的核心公式是:
x_norm = (x - mean) / std
mean:这一行(也就是某个 token 自己的 d_model 维向量)的均值std:这一行的标准差
注意:LayerNorm 是逐行独立计算的——"猫"这一行只看自己的 d_model 个数,不看"吃"和"鱼"那两行,也不看 batch 里其他句子。这也是 LayerNorm 和 BatchNorm 的关键区别,序列长度不固定的 NLP 场景更适合用前者。
用具体数字过一遍(为了计算干净,这里假设 d_model=4,"猫"这个 token 在残差之后的向量是):
x1_猫 = [2.0, -1.0, 3.0, 0.0]
第一步,算均值:
mean = (2.0 + (-1.0) + 3.0 + 0.0) / 4 = 1.0
第二步,减均值(得到"残差",这四个数加起来一定是 0):
[2.0-1.0, -1.0-1.0, 3.0-1.0, 0.0-1.0] = [1.0, -2.0, 2.0, -1.0]
第三步,算方差(平方和再除以 d):
variance = (1.0² + (-2.0)² + 2.0² + (-1.0)²) / 4 = (1+4+4+1)/4 = 2.5
std = √2.5 ≈ 1.581
第四步,除以标准差:
x_norm = [1.0, -2.0, 2.0, -1.0] / 1.581 ≈ [0.632, -1.265, 1.265, -0.632]
验证一下:这四个数加起来约等于 0(均值=0),平方和是 0.4+1.6+1.6+0.4=4.0,除以 d=4 正好等于 1(方差=1)。
一个容易踩的坑:减均值这一步,不是把每个值都"拉向 0"、抹掉信息,而是把这一行数值的参照系挪到了以 0 为中心。相对大小关系(谁大谁小、谁正谁负)基本保留,只是整体的"偏移量"被去掉了。
三、从"方差=1"到几何图形:这一步到底在约束什么?
方差=1 这句话,单独看是个很抽象的统计量,不好想象。但只要把它还原回定义式,就能挖出背后的几何意义。
方差的定义本身就带着"除以 d":
variance = Σ(xᵢ - mean)² / d
LayerNorm 要求 variance=1,而且此时 mean 已经是 0,代入化简:
Σxᵢ² / d = 1 → Σxᵢ² = d
两边开根号:
√(Σxᵢ²) = √d
而 √(Σxᵢ²) 正是这个向量的欧几里得范数(L2 范数 / 模长)——也就是这个点到原点的直线距离,和二维平面上"两点间距离公式"是同一套逻辑,只是从 2 维扩展到了 d 维。
结论:方差=1,等价于"这一行向量到原点的距离被固定为 √d"。
用刚才的例子验证:x_norm ≈ [0.632, -1.265, 1.265, -0.632],模长 = √(0.4+1.6+1.6+0.4) = √4 = 2,而 √d = √4 = 2,完全吻合。
四、可视化:为什么说 LayerNorm 像一个"镂空塑料球"


把上面推导的两个约束叠在一起看:
- 减均值 → 这一行的数字加起来等于 0 → 对应空间里的一个平面
- 除标准差 → 这一行的平方和等于 d → 对应空间里的一个球面(半径 √d)
两个约束同时满足,几何上就是"平面 ∩ 球面"。当 d_model=3 时,这个交集正好是一个可以直接画出来的圆(因为平面正好穿过球心,交出来的是球的"大圆")。
更直观的画法是把每一行向量画成从原点出发的箭头:所有归一化后的向量,箭头长度完全相同(都是 √d),但指向的方向可以是任意角度——方向承载语义信息,长度被统一锁死。这就是为什么可以把 LayerNorm 类比成生活里那种表面全是点、中间镂空的塑料球:所有点都精确落在球的表面,不会跑到球心附近,也不会跑到球外。
d_model 更大(比如项目代码里实际用的 8)时,这个几何结构会升级成更高维的球面,没法直接画出来,但道理完全一样。
需要提醒一句容易混淆的地方:LayerNorm 固定的距离是 √d,不是 1。如果想要模长精确等于 1 的球面,那是另一种操作,叫 L2 归一化(x / ||x||),和 LayerNorm 目的相似但公式不同,别搞混。
五、γ、β:为什么不能只做归一化就完事
完整的 LayerNorm 公式其实还有一步:
LN(x) = γ * x_norm + β
γ 和 β 是两个可学习参数(每个维度一个)。如果只做归一化不加这两个参数,相当于强行把每一层的输出都摁成"方差=1、均值=0"的分布,模型没有任何调节空间。加上 γ、β 之后,模型在训练过程中可以自己学习:这个维度到底要不要保持归一化后的样子,还是应该被放大、缩小或者平移一点。
一句话总结:LayerNorm 不是简单粗暴的"压平",而是"先统一尺度,再把调节权交还给模型"。
六、为什么要做这件事:训练稳定性
数据经过一层一层的 Attention、FFN 之后,每一层输出的数值分布(均值、方差)会漂移——这一层可能普遍偏大,下一层又偏小,层数一深,这种漂移会累积,导致数值要么爆炸要么消失,梯度也会跟着不稳定。
LayerNorm 相当于在每一层入口"清洗"一次分布,把它重新拉回统一、可控的范围。不管前面经过了多少层、数值飘到哪里去了,进入下一个子层的输入始终是干净、稳定的分布——这对深层网络能不能训练起来非常关键。
七、代码实现
免责声明:下面代码里的权重是随机初始化的,数值本身不带任何语义,只是用来验证公式和 shape 是否正确。
import torch
import torch.nn as nn
# ---- 手写实现,验证公式 ----
class MyLayerNorm(nn.Module):
def __init__(self, d_model, eps=1e-5):
super().__init__()
self.gamma = nn.Parameter(torch.ones(d_model)) # [d_model]
self.beta = nn.Parameter(torch.zeros(d_model)) # [d_model]
self.eps = eps
def forward(self, x):
# x: [batch, seq_len, d_model]
mean = x.mean(dim=-1, keepdim=True) # [batch, seq_len, 1]
var = x.var(dim=-1, keepdim=True, unbiased=False) # [batch, seq_len, 1]
x_norm = (x - mean) / torch.sqrt(var + self.eps) # [batch, seq_len, d_model]
return self.gamma * x_norm + self.beta # [batch, seq_len, d_model]
# ---- 用"猫吃鱼"验证:3个token,d_model=8 ----
torch.manual_seed(42)
x1 = torch.randn(1, 3, 8) # [batch=1, seq_len=3, d_model=8],假设这是 EP07 残差之后的输出
my_ln = MyLayerNorm(d_model=8)
out_my = my_ln(x1)
# 官方实现对比(权重同样随机初始化,数值不会完全一致,但公式逻辑一致)
official_ln = nn.LayerNorm(8)
out_official = official_ln(x1)
print("手写实现输出 shape:", out_my.shape) # [1, 3, 8]
print("官方实现输出 shape:", out_official.shape) # [1, 3, 8]
# 验证核心结论:去掉 gamma/beta 之后,每一行的均值应约等于0,方差应约等于1
x_norm_only = (x1 - x1.mean(dim=-1, keepdim=True)) / torch.sqrt(
x1.var(dim=-1, keepdim=True, unbiased=False) + 1e-5
)
print("每行均值(应≈0):", x_norm_only.mean(dim=-1))
print("每行方差(应≈1):", x_norm_only.var(dim=-1, unbiased=False))
# 验证模长:方差=1 意味着模长应该是 sqrt(d_model) = sqrt(8) ≈ 2.828
norm_length = x_norm_only.norm(dim=-1)
print("每行模长(应≈√8≈2.828):", norm_length)
跑一遍这段代码,你会看到"每行方差"确实都逼近 1,"每行模长"确实都逼近 √8 ≈ 2.828——这就是本篇前面几节推导出的几何结论,在代码层面的直接验证。