一言
不管你说再多的慌,只有自己的内心,是无法欺骗的啊。——七大罪
从 LayerNorm 到 RMSNorm:为什么可以去掉均值?

在现代 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,归一化都是用的这个归一化方式。

暂无评论

发送评论 编辑评论

|´・ω・)ノ
ヾ(≧∇≦*)ゝ
(☆ω☆)
(╯‵□′)╯︵┴─┴
 ̄﹃ ̄
(/ω\)
∠( ᐛ 」∠)_
(๑•̀ㅁ•́ฅ)
→_→
୧(๑•̀⌄•́๑)૭
٩(ˊᗜˋ*)و
(ノ°ο°)ノ
(´இ皿இ`)
⌇●﹏●⌇
(ฅ´ω`ฅ)
(╯°A°)╯︵○○○
φ( ̄∇ ̄o)
ヾ(´・ ・`。)ノ"
( ง ᵒ̌皿ᵒ̌)ง⁼³₌₃
(ó﹏ò。)
Σ(っ °Д °;)っ
( ,,´・ω・)ノ"(´っω・`。)
╮(╯▽╰)╭
o(*////▽////*)q
>﹏<
( ๑´•ω•) "(ㆆᴗㆆ)
😂
😀
😅
😊
🙂
🙃
😌
😍
😘
😜
😝
😏
😒
🙄
😳
😡
😔
😫
😱
😭
💩
👻
🙌
🖕
👍
👫
👬
👭
🌚
🌝
🙈
💊
😶
🙏
🍦
🍉
😣
Source: github.com/k4yt3x/flowerhd
颜文字
Emoji
小恐龙
花!
上一篇