在现代 LLM 中,RMSNorm 已经成为非常常见的归一化方式。它可以看作是 LayerNorm 的一个简化版本。
要理解 RMSNorm,最自然的方式是先从 LayerNorm 开始。
1. LayerNorm 在做什么?
对于一个 hidden state:
LayerNorm 的计算公式为:
这里的和是可学习的参数,方便模型不被归一化操作限制,允许模型自己学:每个 feature 最终应该放大还是缩小。
其中:
这个公式主要包含两步:
1.1 减去均值:re-centering
首先:它会把数据的中心移动到 0。
例如:
[99,100,101]
减去均值 100 后:
[-1,0,1]
这样做可以消除整体平移带来的影响。因此 LayerNorm 对整体平移具有不变性:
减去均值会消除hidden state的整体偏移,使归一化结果只与各维度相对于均值的偏离有关,而与整个向量位于什么绝对位置无关。
1.2 除以标准差:re-scaling
接下来:
用于控制 activation 的整体尺度。
例如两个向量:
[1,2,3]
和:
[100,200,300]
虽然数值大小相差 100 倍,但它们具有相同的相对结构。
除以标准差可以消除这种整体尺度变化,避免随着网络不断加深,activation 的 magnitude 变得过大或过小。
这里可以理解为:除以标准差进一步消除整体尺度的影响,使一个向量即使被整体放大或缩小,归一化后的表示仍基本不变。 这个特性叫重缩放不变性(re-scaling invariance)
概括一下就是:
减均值负责控制 center,除标准差负责控制 scale。
2. 从 LayerNorm 到 RMSNorm
RMSNorm(Root Mean Square Layer Normalization)的作者提出了一个很自然的问题:
LayerNorm 中的 re-centering,也就是减均值这一步,真的有必要吗?
实验发现,在很多深度神经网络中,去掉减均值操作后,模型效果几乎没有明显下降。
于是 RMSNorm 直接删除了 re-centering,只保留对尺度的控制。但是之前的除以标准差的操作也是需要均值的,所以作者采用了Root Mean Square作为一组数据scale的衡量,由于 RMS 与向量的 L2 范数成正比,并具有良好的缩放性质,因此可以自然地作为 hidden state 尺度的衡量。
最终:
相比 LayerNorm:
RMSNorm 不再计算均值,也不再进行减均值操作,结构更加简单。
3. 为什么去掉均值还能 work?
一个很重要的直觉是:
归一化最重要的作用之一,是控制 activation 的尺度(magnitude),而不一定需要控制它的位置(center)。
对于:
x=[1,2,3]
如果整体放大 100 倍:
x'=[100,200,300]
虽然数值变大了,但向量内部的相对关系并没有改变。
RMSNorm 会通过 RMS 把这种整体尺度变化消除:
因此,它依然具有很重要的 scale invariance。
从几何角度来看:
可以把 RMSNorm 理解为:
尽量保持 hidden state 的方向不变,只把整个向量的长度调整到一个稳定的尺度。
而 LayerNorm 相比之下,还会额外消除向量整体平移的信息。
RMSNorm 的核心发现就是:
对于模型训练稳定性来说,控制 activation 的 magnitude 往往已经足够,re-centering 带来的额外收益并没有想象中那么重要。
因此可以用更简单的 RMSNorm 代替 LayerNorm,在保持模型效果的同时减少一部分计算。
代码实现
代码实现非常简单:下边copy一下qwen3的源码
class Qwen3NextRMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.zeros(dim))
def _norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
output = self._norm(x.float())
# Llama does x.to(float16) * w whilst Qwen3Next is (x * w).to(float16)
# See https://github.com/huggingface/transformers/pull/29402
output = output * (1.0 + self.weight.float())
return output.type_as(x)
torch.rsqrt对于底层GPU计算是更高效的方式
总结
RMSNorm是一个为数不多没咋变的结构,从遥远的T5时期到现代LLM,归一化都是用的这个归一化方式。









