LLM 网络结构的优化:RMSNorm

1. LLM 的一般网络结构

LLM 是在 Transformer 的基础上,只取其中的 Decoder 部分得到的,LLM 的一般网络结构[1]如下图所示:

可以看到在 LLM 网络结构中使用的是 LayerNorm,Layer Normalization[2] 最初是 2016 年在 RNN 网络中提出并使用,主要解决的问题是在文本的长度都不相同时,无法使用 Batch Normalization,因为 Batch Normalization 要求的是输入的文本长度相同。

2. 从 LayerNorm 到 RMSNorm

2.1. LayerNorm 的基本原理

Layer Normalization[2] 是在 Batch Normalization 的基础上发展起来的,对于神经网络来说,Normalization 的操作是必不可少的,简单来说就是经过多层神经网络的计算后,输出的值会不断膨胀或者缩小,这会导致在反向传播的过程中出现梯度爆炸或者梯度消失,这就需要将数值控制在一定的范围内。首先,我们先看一下 Batch Normalization 如何做归一化的,对于 Batch Normalization,需要分为训练阶段和推理阶段两部分来看,在训练阶段,对于一个包含了 mm 个样本的 batch 来说,B={x1,x2,,xm}B=\left\{ x_1,x_2,\cdots ,x_m\right\},计算的方法分为四步:

  1. 计算批次的均值

μB=1mi=1mxi\mu _B=\frac{1}{m}\sum_{i=1}^{m}x_i

  1. 计算批次的方差

σB2=1mi=1m(xiμB)\sigma _B^2=\frac{1}{m}\sum_{i=1}^{m}\left ( x_i-\mu _B \right )

  1. 归一化

xi^=xiμBσB2+ε\hat{x_i}=\frac{x_i-\mu _B}{\sqrt{\sigma _B^2+\varepsilon }}

其中,ε\varepsilon 是一个极小的常数,作用是防止分母为 0。

  1. 缩放与平移

yi=γxi^+βy_i=\gamma \hat{x_i}+\beta

其中,γ\gamma 称为缩放因子,β\beta 称为偏移因子,均是可训练的参数。

然而,在推理阶段,通常我们都是一次处理一个样本,也就是 m=1m=1 ,因此无法计算批次的均值和方差,因此,需要在训练时维护两个变量,一个被称为全局运行均值 μrunning\mu _{running},一个被称为全局运行方差 σrunning2\sigma _{running}^2,并采用指数移动平均来更新,在训练的每一步的更新公式为:

μrunning=αμrunning+(1α)μB\mu _{running}=\alpha \mu _{running}+\left ( 1-\alpha \right )\mu _B

σrunning2=ασrunning2+(1α)σB2\sigma _{running}^2 =\alpha \sigma _{running}^2+\left ( 1-\alpha \right )\sigma _{B}^2

其中,α\alpha 通常取 0.9。在推理时,直接使用 μrunning\mu _{running}σrunning2\sigma _{running}^2

yi=γxiμrunningσrunning2+ε+βy_i=\gamma \frac{x_i-\mu _{running}}{\sqrt{\sigma _{running}^2+\varepsilon }}+\beta

这里还有一个问题没有解释清楚,那就是对于包含 mm 个样本的 batch 来说,每一个样本究竟是什么样的表示?一般可有下面两种数据格式:

  1. 全连接层

假设全连接层的输出为 [m,d]\left[ m,d\right],其中 mm 是样本的个数,dd 为特征的维度。那么 Batch Normalization 针对的是每一个特征维度计算均值,因此会得到 dd 组的 μ\muσ2\sigma ^2

  1. 卷积层

假设卷积层的输出为 [m,C,H,W]\left[ m,C,H,W \right],其中,mm 是样本的个数,CC 是通道个数,HH 为高,WW 为宽,那么 Batch Normalization 针对的是每一个通道的所有像素点计算均值,因此会得到 CC 组的 μ\muσ2\sigma ^2,此时的均值和方差的分母为 m×H×Wm\times H\times W

而对于文本的数据,输出一般为 [m,n,d]\left[ m,n,d \right],其中,mm 是样本的个数,nn 为序列长度,dd 为特征的维度,如果按照 Batch Normalization 的计算方法,应该针对的是每一个特征维度计算均值,因此会得到 dd 组的 μ\muσ2\sigma ^2,此时的均值和方差的分母为 m×nm\times n

我们注意到一个问题,文本的序列长度 nn 通常都不是固定的,有变长的情况,且在计算均值和方差时,强行变成对 m×nm\times n 求均值,破坏了文本的顺序。

基于以上两点,应该改变计算均值和方差的方法,这便有了 Layer Normalization。Layer Normalization 选择在特征的维度上做均值和方差,假设输出的维度还是 [m,n,d]\left[ m,n,d \right],对于输出的 XRm×n×dX\in \mathbb{R}^{m\times n\times d},我们需要对第 ii 个样本的第 tt 个 token 分别计算:

  1. 计算均值:

μi,t=1dk=1dXi,t,k\mu _{i,t}=\frac{1}{d}\sum_{k=1}^{d}X_{i,t,k}

  1. 计算方差:

σi,t2=1di=1d(Xi,t,kμi,t)\sigma _{i,t}^2=\frac{1}{d}\sum_{i=1}^{d}\left ( X_{i,t,k}-\mu _{i,t} \right )

  1. 归一化

Xi,t,k^=Xi,t,kμi,tσi,t2+ε\hat{X_{i,t,k}}=\frac{X_{i,t,k}-\mu _{i,t}}{\sqrt{\sigma _{i,t}^2+\varepsilon }}

  1. 缩放与平移

Yi,t,k=γkXi,t,k^+βkY_{i,t,k}=\gamma _k\hat{X_{i,t,k}}+\beta _k

与 Batch Normalization 中不一样的是,在 Layer Normalization 中 γk\gamma _kβk\beta _k 都是 dd 维的。另一点是在 Layer Normalization 不再区分训练阶段和推理阶段,两个阶段的计算逻辑完全一致。

2.2. RMSNorm 的基本原理

RMSNorm(Root Mean Square Normalization)[3]是目前大模型中最流行的归一化方法,实际上是 Layer Normalization 的简化版本。在 RMSNorm 中,仅通过计算均方根(RMS)来进行归一化,无需再计算均值,同时也去除掉了 βk\beta _k。其计算公式如下:

RMS(Xi,t,k)=γkXi,t,k1di=1d(Xi,t,k)+εRMS\left( X_{i,t,k} \right)=\gamma _k\frac{X_{i,t,k}}{\sqrt{\frac{1}{d}\sum_{i=1}^{d}\left ( X_{i,t,k} \right )+\varepsilon }}

在参考文献[3]中的实验结果显示这种简化并没有对模型的性能产生明显影响,反而能提升训练的速度。我们参考 Llama 的实现[4]

class RMSNorm(torch.nn.Module):
    def __init__(self, dim: int, eps: float = 1e-6):
        """
        Initialize the RMSNorm normalization layer.

        Args:
            dim (int): The dimension of the input tensor.
            eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.

        Attributes:
            eps (float): A small value added to the denominator for numerical stability.
            weight (nn.Parameter): Learnable scaling parameter.

        """
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim))

    def _norm(self, x):
        """
        Apply the RMSNorm normalization to the input tensor.

        Args:
            x (torch.Tensor): The input tensor.

        Returns:
            torch.Tensor: The normalized tensor.

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

    def forward(self, x):
        """
        Forward pass through the RMSNorm layer.

        Args:
            x (torch.Tensor): The input tensor.

        Returns:
            torch.Tensor: The output tensor after applying RMSNorm.

        """
        output = self._norm(x.float()).type_as(x)
        return output * self.weight

3. Post Norm 和 Pre Norm

关于归一化,除了上述不同归一化的影响外,针对归一化在大模型中相对于残差连接的位置,也有 Post Norm 和 Pre Norm 之分[5],如下图所示:

Post Norm 虽然理论优美,能够保证输出分布稳定,但在深层网络中极难训练,而 Pre Norm 彻底解决了这个问题。Pre Norm 确保了残差连接的主干道上没有任何线性变换或非线性层,保留了完美的恒等映射。同时,网络越深,这种“直通车道”的优越性越明显。

4. 总结

Batch Norm 的归一化方法不适合在文本模型上做归一化,最重要的原因是文本长度并不固定,同时,基于 Batch Norm 的方法会破坏文本 token 之间的相对顺序,因此才有了 Layer Norm,同时,为了简化,又可以省略掉均值,因此就有了 RMS Norm,针对 Norm 在网络中的位置,Pre Norm 相比较与 Post Norm 来说,使得网络更加容易训练。

参考文献

[1] LLM 推理优化技术:KV Cache

[2] Ba J L, Kiros J R, Hinton G E. Layer normalization[J]. arXiv preprint arXiv:1607.06450, 2016.

[3] Zhang B, Sennrich R. Root mean square layer normalization[J]. Advances in neural information processing systems, 2019, 32.

[4] https://github.com/meta-llama/llama/blob/main/llama/model.py

[5] Xiong R, Yang Y, He D, et al. On layer normalization in the transformer architecture[C]//International conference on machine learning. PMLR, 2020: 10524-10533.