RMSNorm,它是很多现代大模型(包括 DeepSeek 系列)在 LayerNorm 的基础上改进的一种归一化方法,主要目的是提高计算效率和稳定性。

1. 背景

在 Transformer 中,常见的归一化方法是 LayerNorm
LayerNorm(x)=x−μσ⋅γ+β \text{LayerNorm}(x) = \frac{x - \mu}{\sigma} \cdot \gamma + \beta LayerNorm(x)=σxμγ+β
其中:

  • μ=mean(x)\mu = \text{mean}(x)μ=mean(x)
  • σ=Var(x)\sigma = \sqrt{\text{Var}(x)}σ=Var(x)
  • γ,β\gamma, \betaγ,β 是可学习的缩放和平移参数

LayerNorm 的问题:

  • 需要计算均值和方差,这两个操作都涉及减法和除法,开销较大。
  • 方差计算要遍历整个向量,增加了一些数值不稳定性。

2. RMSNorm 的思想

RMSNorm(Root Mean Square Layer Normalization)去掉了 减均值 这一步,只使用 均方根(Root Mean Square, RMS)来归一化。

定义:
RMSNorm(x)=xRMS(x)⋅γ \text{RMSNorm}(x) = \frac{x}{\text{RMS}(x)} \cdot \gamma RMSNorm(x)=RMS(x)xγ
其中:
RMS(x)=1n∑i=1nxi2+ϵ \text{RMS}(x) = \sqrt{\frac{1}{n}\sum_{i=1}^n x_i^2 + \epsilon} RMS(x)=n1i=1nxi2+ϵ

  • nnn 是向量维度
  • ϵ\epsilonϵ 是防止除零的小数
  • γ\gammaγ 是可训练的缩放参数
  • 没有 β\betaβ(可加,但很多实现省略)

与 LayerNorm 对比:

特性 LayerNorm RMSNorm
去均值 ✅ 有减去均值 ❌ 无减均值
方差计算 需要 不需要(只计算均方根)
参数 γ,β\gamma, \betaγ,β γ\gammaγ(可选 β\betaβ
数值稳定性 计算更复杂,可能受均值影响 更简单、更稳定
计算性能 稍慢 更快(少一步减均值和方差计算)

3. 为什么大模型喜欢用 RMSNorm

  1. 减少计算量:去掉均值计算和方差计算,GPU 上更高效。
  2. 数值稳定性好:只依赖模长,不受均值漂移影响。
  3. 梯度更稳定:特别是在超大批量和长序列下,减少归一化带来的噪声。
  4. 工程简单化:尤其是在混合精度(fp16/bf16)训练时,减少精度丢失风险。

4. 应用场景

  • Transformer 架构中替代 LayerNorm(例如 Pre-LN Transformer 可以改成 Pre-RMSNorm)。
  • LLM(大语言模型):GPT-NeoX、LLaMA、DeepSeek 都有用 RMSNorm。
  • 底层块的激活归一化,尤其在 MoE 中,减少 LayerNorm 复杂度。

5. 公式总结

假设输入为向量 x∈Rnx \in \mathbb{R}^nxRn
RMSNorm 的计算过程:
RMSNorm(x)=x1n∑i=1nxi2+ϵ⋅γ \text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{n} \sum_{i=1}^n x_i^2 + \epsilon}} \cdot \gamma RMSNorm(x)=n1i=1nxi2+ϵ xγ


💡 小结:
RMSNorm 是 更轻量的 LayerNorm,用均方根代替标准差,去掉均值归一化,既快又稳定,特别适合大模型(比如 DeepSeek V3)这样算力密集的场景。

用一个具体的向量例子一步步演示 RMSNorm 的计算过程直观理解和 LayerNorm 的区别。


例子

假设我们有一个输入向量:
x=[2.0,−1.0,3.0] x = [2.0, -1.0, 3.0] x=[2.0,1.0,3.0]
维度 n=3n = 3n=3,缩放参数 γ=[1.5,1.5,1.5]\gamma = [1.5, 1.5, 1.5]γ=[1.5,1.5,1.5](每个维度都有一个可学习的比例),ϵ=10−8\epsilon = 10^{-8}ϵ=108


Step 1:计算均方根 RMS

RMS 的定义:
RMS(x)=1n∑i=1nxi2+ϵ \text{RMS}(x) = \sqrt{\frac{1}{n} \sum_{i=1}^n x_i^2 + \epsilon} RMS(x)=n1i=1nxi2+ϵ

先算平方:
[2.02,(−1.0)2,3.02]=[4.0,1.0,9.0] [2.0^2, (-1.0)^2, 3.0^2] = [4.0, 1.0, 9.0] [2.02,(1.0)2,3.02]=[4.0,1.0,9.0]

求和:
4.0+1.0+9.0=14.0 4.0 + 1.0 + 9.0 = 14.0 4.0+1.0+9.0=14.0

求平均:
14.03≈4.6667 \frac{14.0}{3} \approx 4.6667 314.04.6667

开方:
4.6667+10−8≈2.1602 \sqrt{4.6667 + 10^{-8}} \approx 2.1602 4.6667+108 2.1602

所以:
RMS(x)≈2.1602 \text{RMS}(x) \approx 2.1602 RMS(x)2.1602


Step 2:归一化

x^=xRMS(x) \hat{x} = \frac{x}{\text{RMS}(x)} x^=RMS(x)x
计算:
[2.0,−1.0,3.0]÷2.1602≈[0.9258,−0.4629,1.3887] [2.0, -1.0, 3.0] \div 2.1602 \approx [0.9258, -0.4629, 1.3887] [2.0,1.0,3.0]÷2.1602[0.9258,0.4629,1.3887]


Step 3:乘以缩放系数 γ\gammaγ

y=x^⋅γ y = \hat{x} \cdot \gamma y=x^γ
如果 γ=[1.5,1.5,1.5]\gamma = [1.5, 1.5, 1.5]γ=[1.5,1.5,1.5]
y≈[0.9258×1.5,−0.4629×1.5,1.3887×1.5] y \approx [0.9258 \times 1.5, -0.4629 \times 1.5, 1.3887 \times 1.5] y[0.9258×1.5,0.4629×1.5,1.3887×1.5]
y≈[1.3887,−0.6944,2.0831] y \approx [1.3887, -0.6944, 2.0831] y[1.3887,0.6944,2.0831]


最终输出

RMSNorm(x)≈[1.3887,−0.6944,2.0831] \text{RMSNorm}(x) \approx [1.3887, -0.6944, 2.0831] RMSNorm(x)[1.3887,0.6944,2.0831]


和 LayerNorm 对比

为了对比,如果是 LayerNorm:

  1. 先减去均值:
    μ=2.0−1.0+3.03=1.3333 \mu = \frac{2.0 - 1.0 + 3.0}{3} = 1.3333 μ=32.01.0+3.0=1.3333
    得到:
    x−μ≈[0.6667,−2.3333,1.6667] x - \mu \approx [0.6667, -2.3333, 1.6667] xμ[0.6667,2.3333,1.6667]

  2. 再除以标准差:
    σ=(0.6667)2+(−2.3333)2+(1.6667)23≈1.6997 \sigma = \sqrt{\frac{(0.6667)^2 + (-2.3333)^2 + (1.6667)^2}{3}} \approx 1.6997 σ=3(0.6667)2+(2.3333)2+(1.6667)2 1.6997
    除后:
    [0.3920,−1.3728,0.9800] [0.3920, -1.3728, 0.9800] [0.3920,1.3728,0.9800]
    再乘 γ\gammaγ 得到输出。

可以看到:

  • RMSNorm 直接用模长来缩放,不减均值,数值变化更简单。
  • LayerNorm 会调整中心位置(均值归零),对数据分布影响更大。
Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐