代码实现 RMSNorm

从零实现 RMSNorm:一行公式如何变成 PyTorch 代码

先回忆下 RMSNorm(Root Mean Square Normalization)的原理:

RMS(x)=1di=1dxi2

以及 RMSNorm 的计算公式:

y=xRMS(x)γ

其中:

接下来将结合一段真实的大模型代码,逐行讲解 RMSNorm 的实现原理。


一、完整代码

class RMSNorm(torch.nn.Module):
    def __init__(self, dim: int, eps: float = 1e-5):
        super().__init__()
        self.eps = eps
        self.weight = torch.nn.Parameter(torch.ones(dim))

    def norm(self, x):
        return x * torch.rsqrt(
            x.pow(2).mean(-1, keepdim=True) + self.eps
        )

    def forward(self, x):
        return (self.weight * self.norm(x.float())).type_as(x)

这段代码非常短。

真正的核心逻辑只有一行:

x * torch.rsqrt(
    x.pow(2).mean(-1, keepdim=True) + self.eps
)

理解这一行,基本就理解了 RMSNorm。


二、继承 torch.nn.Module

class RMSNorm(torch.nn.Module):

PyTorch 中所有神经网络层都继承自:

torch.nn.Module

例如:

nn.Linear
nn.Conv2d
nn.LayerNorm

本质上也都是 Module。

这样做的好处是:

例如:

model = RMSNorm(512)

for name, param in model.named_parameters():
    print(name)

输出:

weight

说明 RMSNorm 中存在一个可学习参数。


三、初始化函数

init()

def __init__(self, dim: int, eps: float = 1e-5):

参数:

参数 含义
dim 隐藏层维度
eps 防止除零的小常数

例如:

RMSNorm(4096)

表示:

hidden_size = 4096

这是 LLaMA、Qwen 等模型中常见的隐藏维度。


四、为什么要有 eps?

self.eps = eps

默认:

eps = 1e-5

原因很简单。

假设输入:

x = [0,0,0]

那么:

RMS(x)=0

此时会出现:

x0

程序直接崩溃。

因此需要:

RMS(x) + eps

变成:

mean(x2)+105

避免除零错误。


五、可学习参数 weight

torch.ones

torch.ones(dim)

例如:

torch.ones(4)

输出:

tensor([1., 1., 1., 1.])

表示创建一个全是 1 的向量。


torch.nn.Parameter

self.weight = torch.nn.Parameter(
    torch.ones(dim)
)

等价于:

γ = [1,1,1,...]

这正对应 RMSNorm 公式中的:

γ

为什么需要 weight?

如果没有 weight:

y=xRMS(x)

模型表达能力会下降。

因此加入:

y=γxRMS(x)

让模型自己学习:

这就是:

self.weight

的作用。


六、进入核心:norm()

def norm(self, x):

这里真正实现 RMSNorm。


第一步:平方

x.pow(2)

例如:

x = [2,4,6]

得到:

[4,16,36]

对应数学公式:

xi2

第二步:求均值

.mean(-1, keepdim=True)

继续上面的例子:

[4,16,36]

均值:

4+16+363=18.67

得到:

18.67

对应公式:

1dxi2

为什么是 -1?

mean(-1)

表示:

最后一个维度

例如:

x.shape
=
(2,3)
[
 [1,2,3],
 [4,5,6]
]

执行:

x.mean(-1)

得到:

[
 2,
 5
]

即:

每个 token 单独计算 RMS。


keepdim=True

如果:

x.shape=(2,3)

那么:

x.mean(-1)

结果:

shape=(2,)

变成:

[2,5]

而:

x.mean(-1, keepdim=True)

结果:

shape=(2,1)

即:

[
 [2],
 [5]
]

这样方便后续广播运算(Broadcast)。


七、加上 eps

x.pow(2).mean(-1, keepdim=True)
+ self.eps

对应:

1dxi2+ϵ

避免分母为 0。


八、torch.rsqrt()

这是 RMSNorm 中最容易被忽略的一步。


rsqrt 是什么?

torch.rsqrt(x)

表示:

1x

例如:

torch.rsqrt(
    torch.tensor([1.,4.,9.])
)

输出:

[1.0,0.5,0.3333]

因为:

11=114=0.519=0.333

为什么不用 sqrt?

RMSNorm 需要:

x1dxi2+ϵ

而:

torch.rsqrt(...)

直接得到:

11dxi2+ϵ

于是:

x * rsqrt(...)

就等价于:

x1dxi2+ϵ

少做一次除法。

GPU 上更高效。


九、完成 RMS 归一化

这一句:

x * torch.rsqrt(
    x.pow(2).mean(-1, keepdim=True)
    + self.eps
)

完整对应:

x1dxi2+ϵ

这就是 RMSNorm 的核心数学公式。


十、forward()

def forward(self, x):

PyTorch 约定:

output = model(x)

实际上调用:

model.forward(x)

x.float()

self.norm(x.float())

很多大模型使用:

float16
bfloat16

训练。

例如:

dtype=torch.float16

为了提高数值稳定性:

归一化时先转成:

float32

即:

x.float()

这样:

RMS

计算更加准确。


乘以 weight

self.weight * self.norm(...)

对应数学公式:

y=γxRMS(x)

其中:

γ=self.weight

type_as(x)

.type_as(x)

表示:

结果转回原来的数据类型

例如:

输入:

float16

归一化时:

float32

最终再转回:

float16

这样:

这是现代 LLM 的标准做法。


十一、完整执行流程

假设输入:

x = [2,4,6]

Step1:平方

[4,16,36]

Step2:求均值

18.67

Step3:加 eps

18.67001

Step4:rsqrt

1 / sqrt(18.67001) ≈ 0.231

Step5:乘回原向量

[
 0.462,
 0.924,
 1.386
]

Step6:乘 weight

假设:

weight = [1,1,1]

结果不变:

[
 0.462,
 0.924,
 1.386
]

如果训练后:

weight =
[
 1.2,
 0.8,
 1.5
]

结果变成:

[
 0.554,
 0.739,
 2.079
]

模型便可以自主调整每个维度的重要性。


十二、一句话总结

RMSNorm 的实现本质上只有一句核心代码:

x * torch.rsqrt(
    x.pow(2).mean(-1, keepdim=True) + eps
)

它对应数学公式:

x1dxi2+ϵ

然后再乘上可学习参数:

y=γx1dxi2+ϵ

相比 LayerNorm,RMSNorm 去掉了均值计算和方差计算,只保留尺度归一化(Scaling),因此实现更简单、计算更高效,也成为 LLaMA、Qwen、DeepSeek 等现代大模型的默认归一化方案。