← 返回文章档案

Attention Residuals 阅读补充

围绕 Attention Residuals 逐步整理相关知识、数学推导与系统分析,先从 RMSNorm 的数学原理及其与 LayerNorm 的差异开始。

这篇文章用于沉淀和分析阅读苏神(苏剑林)《Attention Residuals 回忆录》 时涉及的相关知识。内容会随阅读逐步增加,每一章集中处理一个需要补充推导或实现背景的部分。

RMSNorm 的数学原理

RMSNorm 用输入向量在特征维度上的二阶原点矩衡量整体幅度,再按这个幅度缩放整个向量。

设一个 token 的隐藏状态为

x=(x1,x2,,xd)Rd.\boldsymbol{x}=(x_1,x_2,\ldots,x_d)\in\mathbb{R}^d.

先定义带数值稳定项的均方根

rε(x)=1di=1dxi2+ε,r_\varepsilon(\boldsymbol{x}) =\sqrt{\frac{1}{d}\sum_{i=1}^d x_i^2+\varepsilon},

再对所有维度使用同一个标量完成归一化,并加入可学习的逐维缩放参数 γ\boldsymbol{\gamma}

RMSNorm(x)=γxrε(x).\operatorname{RMSNorm}(\boldsymbol{x}) =\boldsymbol{\gamma}\odot \frac{\boldsymbol{x}}{r_\varepsilon(\boldsymbol{x})}.

这里的均值发生在同一个 token 的 dd 个特征维度上。它是对当前向量的确定性计算,无需估计数据集或 batch 的统计量。RMSNorm 原论文 将其作为 LayerNorm 的简化形式,保留重缩放不变性并省去均值中心化。

从均方根到欧氏范数

暂时令 ε=0\varepsilon=0,并假设 x0\boldsymbol{x}\neq\boldsymbol{0}。均方根可以直接写成欧氏范数:

r0(x)=1di=1dxi2=x2d.r_0(\boldsymbol{x}) =\sqrt{\frac{1}{d}\sum_{i=1}^d x_i^2} =\frac{\lVert\boldsymbol{x}\rVert_2}{\sqrt d}.

记不含可学习参数的归一化结果为 x~\tilde{\boldsymbol{x}},则

x~=xr0(x)=dxx2.\tilde{\boldsymbol{x}} =\frac{\boldsymbol{x}}{r_0(\boldsymbol{x})} =\sqrt d\frac{\boldsymbol{x}}{\lVert\boldsymbol{x}\rVert_2}.

x/x2\boldsymbol{x}/\lVert\boldsymbol{x}\rVert_2 只保留输入方向,前面的 d\sqrt d 将长度设为固定值。直接计算可得

x~22=i=1dxi21dj=1dxj2=d,\lVert\tilde{\boldsymbol{x}}\rVert_2^2 =\frac{\sum_{i=1}^d x_i^2} {\frac{1}{d}\sum_{j=1}^d x_j^2} =d,

因此

x~2=d.\lVert\tilde{\boldsymbol{x}}\rVert_2=\sqrt d.

ε=0\varepsilon=0 的理想条件下,RMSNorm 的归一化部分沿径向把非零输入映射到半径为 d\sqrt d 的超球面。这个映射保留方向并移除长度。

目标长度取 d\sqrt d,等价于让归一化结果的平均平方值为 11

1di=1dx~i2=1.\frac{1}{d}\sum_{i=1}^d \tilde{x}_i^2=1.

因此 RMSNorm 控制的是整个隐藏状态的平均平方幅度。单个坐标仍然可以远大于 11,各个维度也不会分别获得单位方差。

二阶原点矩与方差

RMSNorm 和 LayerNorm 的差异可以从二阶原点矩与方差的关系看出。定义当前向量在特征维度上的均值和方差:

μ=1di=1dxi,σ2=1di=1d(xiμ)2.\mu=\frac{1}{d}\sum_{i=1}^d x_i, \qquad \sigma^2=\frac{1}{d}\sum_{i=1}^d(x_i-\mu)^2.

二阶原点矩为

m2=1di=1dxi2.m_2=\frac{1}{d}\sum_{i=1}^d x_i^2.

展开方差后得到

m2=σ2+μ2.m_2=\sigma^2+\mu^2.

RMSNorm 使用 m2m_2,所以均值偏移 μ2\mu^2 也会进入归一化尺度。LayerNorm 先减去均值,再使用中心二阶矩 σ2\sigma^2。忽略可学习参数时,两者分别为

x~RMS=xσ2+μ2+ε,\widetilde{\boldsymbol{x}}_{\mathrm{RMS}} =\frac{\boldsymbol{x}}{\sqrt{\sigma^2+\mu^2+\varepsilon}}, x~LN=xμ1σ2+ε.\widetilde{\boldsymbol{x}}_{\mathrm{LN}} =\frac{\boldsymbol{x}-\mu\boldsymbol{1}} {\sqrt{\sigma^2+\varepsilon}}.

μ=0\mu=0、两者采用相同的 ε\varepsilon 和缩放参数,并且 LayerNorm 的偏置为零时,两种归一化给出相同结果。当 μ\mu 只接近零时,它们使用的尺度接近,RMSNorm 仍会保留沿 1\boldsymbol{1} 方向的分量。

这个差异也有直接的几何表示。令

P=I1d11T,\boldsymbol{P} =\boldsymbol{I}-\frac{1}{d}\boldsymbol{1}\boldsymbol{1}^{\mathsf T},

Px=xμ1\boldsymbol{P}\boldsymbol{x}=\boldsymbol{x}-\mu\boldsymbol{1}。LayerNorm 先把输入投影到与 1\boldsymbol{1} 正交的均值为零子空间,再在该子空间内归一化;RMSNorm 直接在完整的 Rd\mathbb{R}^d 中按长度归一化。LayerNorm 原论文 给出的定义同时包含中心化、缩放以及归一化后的可学习增益和偏置。

性质RMSNormLayerNorm
归一化统计量二阶原点矩 m2m_2中心二阶矩 σ2\sigma^2
是否减去特征均值
正比例尺度变化ε=0\varepsilon=0 时消除ε=0\varepsilon=0 时消除
整体平移 x+c1\boldsymbol{x}+c\boldsymbol{1}会改变结果归一化结果不变
几何作用在完整空间中按半径归一化先进入均值为零子空间,再按半径归一化

尺度变化与平移

尺度不敏感性是 RMSNorm 在 AttnRes 推导中承担的关键性质。令 a>0a>0,则

rε(ax)=a2m2+ε=am2+εa2.r_\varepsilon(a\boldsymbol{x}) =\sqrt{a^2m_2+\varepsilon} =a\sqrt{m_2+\frac{\varepsilon}{a^2}}.

ε=0\varepsilon=0 时,正比例缩放会被精确消除:

RMSNorm(ax)=RMSNorm(x),a>0.\operatorname{RMSNorm}(a\boldsymbol{x}) =\operatorname{RMSNorm}(\boldsymbol{x}), \qquad a>0.

实际实现采用 ε>0\varepsilon>0,此时等式只近似成立。当 a2m2εa^2m_2\gg\varepsilon 时,稳定项的影响很小。对于 a<0a<0,在 ε=0\varepsilon=0 时有

RMSNorm(ax)=RMSNorm(x),\operatorname{RMSNorm}(a\boldsymbol{x}) =-\operatorname{RMSNorm}(\boldsymbol{x}),

因为输入方向发生了符号翻转。

RMSNorm 对整体平移没有这项不变性。令 x=x+c1\boldsymbol{x}'=\boldsymbol{x}+c\boldsymbol{1},其二阶原点矩变为

1di=1d(xi+c)2=m2+2cμ+c2.\frac{1}{d}\sum_{i=1}^d(x_i+c)^2 =m_2+2c\mu+c^2.

LayerNorm 会在中心化时消除新增的 c1c\boldsymbol{1},所以整体平移不会改变其归一化结果。RMSNorm 保留这部分均值信息,并让它参与尺度计算。

梯度沿切向传播

为了观察归一化怎样改变梯度,先去掉 γ\boldsymbol{\gamma},记

r=1dxTx+ε,y=xr.r=\sqrt{\frac{1}{d}\boldsymbol{x}^{\mathsf T}\boldsymbol{x}+\varepsilon}, \qquad \boldsymbol{y}=\frac{\boldsymbol{x}}{r}.

它的 Jacobian 为

yx=1rI1dr3xxT.\frac{\partial\boldsymbol{y}}{\partial\boldsymbol{x}} =\frac{1}{r}\boldsymbol{I} -\frac{1}{dr^3}\boldsymbol{x}\boldsymbol{x}^{\mathsf T}.

如果上游梯度为 g=L/y\boldsymbol{g}=\partial\mathcal{L}/\partial\boldsymbol{y},输入梯度就是

Lx=1r(gyyTgd).\frac{\partial\mathcal{L}}{\partial\boldsymbol{x}} =\frac{1}{r}\left( \boldsymbol{g} -\boldsymbol{y}\frac{\boldsymbol{y}^{\mathsf T}\boldsymbol{g}}{d} \right).

第一项按输入 RMS 统一调整梯度尺度,第二项减去与输入径向方向相关的分量。当 ε=0\varepsilon=0 时,Jacobian 在径向上的作用为

yxx=0.\frac{\partial\boldsymbol{y}}{\partial\boldsymbol{x}}\boldsymbol{x} =\boldsymbol{0}.

这与正比例尺度不变性一致:沿 x\boldsymbol{x} 方向只改变长度,归一化结果保持不变。对任意满足 xTv=0\boldsymbol{x}^{\mathsf T}\boldsymbol{v}=0 的切向量 v\boldsymbol{v},Jacobian 的作用为 v/r\boldsymbol{v}/r。因此在理想条件下,归一化核心保留切向变化并消除径向变化。

加入 ε\varepsilon 后,径向特征值变为 ε/r3\varepsilon/r^3,径向梯度会被显著减小,但不会严格归零。加入 γ\boldsymbol{\gamma} 后,上述分析仍适用于归一化核心;反向传播时,上游梯度会先逐维乘以 γ\boldsymbol{\gamma}

ε 与 γ 改变了什么

ε\varepsilon 在输入接近零向量时保持分母为正。它也会让输出长度略低于 d\sqrt d

x~22=x22x22/d+ε=dx22x22+dε<d.\lVert\tilde{\boldsymbol{x}}\rVert_2^2 =\frac{\lVert\boldsymbol{x}\rVert_2^2} {\lVert\boldsymbol{x}\rVert_2^2/d+\varepsilon} =\frac{d\lVert\boldsymbol{x}\rVert_2^2} {\lVert\boldsymbol{x}\rVert_2^2+d\varepsilon} <d.

只有在 x22dε\lVert\boldsymbol{x}\rVert_2^2\gg d\varepsilon 时,输出长度才接近 d\sqrt d

γ\boldsymbol{\gamma} 为每个特征维度恢复可学习的尺度。归一化核心在理想条件下把输入映射到超球面,逐维乘以非均匀的 γ\boldsymbol{\gamma} 后,这个超球面的像成为轴对齐的椭球面。模型由此可以调整各个特征维度的有效范围。

这些性质共同提供整体尺度控制:后续层看到的输入 RMS 保持在相对稳定的范围,输入的全局放大对归一化结果影响较小,径向梯度也受到抑制。它们不保证隐藏状态均值为零、不保证每个维度具有单位方差,也不保证 Pre-Norm 架构中的残差流本身具有有界范数。

回到 AttnRes

原文在层间注意部分使用了 RMSNorm 的正比例尺度不变性。假设层函数可以写成

f(z)=F(RMSNorm(z)),\boldsymbol{f}(\boldsymbol{z}) =\boldsymbol{F}(\operatorname{RMSNorm}(\boldsymbol{z})),

那么在 ε=0\varepsilon=0c>0c>0 的条件下有

f(cz)=f(z).\boldsymbol{f}(c\boldsymbol{z})=\boldsymbol{f}(\boldsymbol{z}).

对于一组非负权重 bsb_s,令 B=sbs>0B=\sum_s b_s>0,则

sbsys=BsbsBys.\sum_s b_s\boldsymbol{y}_s =B\sum_s\frac{b_s}{B}\boldsymbol{y}_s.

进入 In Norm 后,外部的正标量 BB 会被消除。因此,把权重归一化到 sas=1\sum_s a_s=1 不会改变后续层接收到的归一化方向;实际实现中的 ε\varepsilon 使该结论成为高信号幅度下的近似。

AttnRes 还使用

at+1,sexp(wt+1TRMSNorm(ys))a_{t+1,s}\propto \exp\left( \boldsymbol{w}_{t+1}^{\mathsf T} \operatorname{RMSNorm}(\boldsymbol{y}_s) \right)

计算层间注意力。这里 RMSNorm 在计算相似度前控制 Key 的整体尺度,使注意力分数主要取决于归一化后的方向和可学习 Query wt+1\boldsymbol{w}_{t+1}。这两个位置分别对应 RMSNorm 的尺度不敏感性和方向保留性质,也是后续理解 AttnRes 数学形式与系统实现的基础。