如何稳定神经网络的前向和反向传播
TODO
核心概览
“优化”问题在整个学习中的定位是什么?
它和具体的上层应用关系不是很大,我们根据具体的问题,比如语言生成,语音识别,图片生成建立起了模型,比如最大化条件概率似然,并在这种建模框架下找到了符合这种模态的可行的神经网络架构(transformer, CNN, 图网络,diffusion),然后就用理论上几乎必然能找到极小值的随机梯度下降进行优化。
但现实中,前向传播和反向传播的梯度都是层层计算而产生的,其中可能过大或过小,因此梯度下降很可能无法真正实施起来,并且随着 scaling 会越来越难控制,这时候数值稳定性问题就出现了,我们的目的是,保证前向和
利用梯度这种局部信息,使得(任意规模的)模型高效稳定地找到特定目标下可行的(能完成特定任务)解。
所以优化问题面对的对象不是应用问题本身(更不是智能或者 AGI),而是模型的架构(层的宽度,深度,激活函数),硬件,数据的表征(浮点数形式)等,数值稳定的核心是各种层面对特征进行尺度的归一化。
我们可以通过 W 的分布去控制方差,也可以通过 W 的奇异值去控制,这两种控制手段分别出现在不同场景,控制分布是在初始化的时候,控制奇异值则是在梯度更新的时候,因为梯度更新时我们要追求效率最高且稳定,如果对梯度进行某种随机化更难以达到高效。
对于向量,稳定性的核心指标是方差,在零均值情况下被称为 squared RMS norm, 它就等于 L2 norm 除以向量长度。
随机初始化控制训练早期阶段特征的数值稳定性,而对梯度下降的改进是控制中后期数值的持续稳定性
Wx 线性变换的稳定性
先考虑神经网络中最基础的模块,向量和矩阵乘法 y = Wx 。
为了保证数值的稳定性,即输入一个数值稳定的 x 经变换后 y 的值也尽可能保持稳定,需要用数学方式来刻画这种稳定性, 而最常见的是向量的 L2 范数和方差(期望为 0 时则称为 squared RMS norm),具体细节在各个小节中展开。
人为能控制 y=Wx 的哪些部分?
如果 y=Wx 中没有一个向量是能够人为修改的,那么控制稳定性只是空谈。
要能保证稳定,意味着能用一些方式去控制其中的某些部分,假设训练时用 batch=1 ,即输入只有一个样本 x, 列举出一般网络里 y=Wx 的常见场景和控制手段:
x 是网络输入的第一层,我们会对完整输入 X 的各个特征维度跨样本地进行标准化成均值为 0 ,方差为 1 的矩阵。 这种情况下单个样本 x 的特征向量各个值的方差自然也会相对稳定。
此外也有对 x 样本内的所有特征进行某种归一化,比如图片输入里所有像素值从 0 到 255 缩小到 0-1 之间,这使得各个样本 x 向量的 norm 或者方差都在一个集中的稳定范围内。
- 前向传播里,W 是全连接层的权重, x 是该层输入(不限于第一层)。 在训练前,一般用随机化初始化方式,能控制的是它的分布、均值和方差,因此,我们要考虑的是随机化的分布,均值和方差如何影响向量的 L2 范数合法方差。
- 反向传播里 W 是局部雅克比矩阵(或其等效形式),比如 \( h_2 = W h_1 \) 层,上游对 h2 的梯度如果是 g, 那么我们会用 W^Tg 计算出损失对 h2 的梯度并传播下去,为了使得梯度不爆炸也不消失,我们能做的有:
- 保证最上游(loss 函数)的梯度数值适中,但我们无法对最上游梯度像第一层输入数据那样进行标准化,我们能控制的是用更平滑地不容易有跳跃或平坦区的损失函数,。例如 SVM 损失存在一些截断区,会使梯度变为 0,从而影响方差的稳定,那么就更少使用。
- 在第一次反向传播的时候 \( W^Tg \) 中的 W 仍然是我们随机初始化的权重,因此随机初始化本身可以同时影响前向和后向传播。
- 之后 W 会因为梯度下降而被小步更新,它会逐渐偏离初始的那个分布,但是我们为了保证它不会影响数值稳定,可以通过对原始梯度进行某种修正(比如乘以足够小的步长,裁剪,权重衰减等)方式使得权重在不断经过叠加改变(梯度下降仅仅是不断加上新的权重矩阵)的过程中,W 的值自身不会太小或太大,并且它对输入(无论是前向特征还是后向上游梯度)的变换也不会导致输出数值不稳定。
Wx 对 L2 范数影响
L2 范数度量向量的整体长度,它是 \(\sqrt{\sum x_i^2}\)。如果这个值很大,要么其中有某些分量很大,要么各个值都相对较大,对 L2 范数施加一个上限约束,比如 1,最极端情况下是只有一个元素为 1 其他为 0 ,相当于要求所有分量都不能大于 1, 因此它某种程度上是对元素最大绝对值的一种软约束。
反过来,它也不能太小,否则所有值都会非常接近于零,很容易引发浮点误差,或者在反向传播时梯度消失。因此一般会希望范数在一个稳定的区间里,比如期望它就等于 1。
接着看 Wx 对 L2 范数的影响:
如果把 W 用 SVD 分解,得到 \ (W = UΣ V^T \),这里 U 和 V 都是单位正交矩阵,在二维几何里对应旋转或反射,因而完全不会改变向量的长度:\(\|V^T x\| = \|x\|\),\( \|U(\Sigma V^T x)\| = \|\Sigma V^T x\| \)。
只是中间的对角矩阵 \(\Sigma\) 会对范数缩放,其对角线上依次排列着从大到小的奇异值 \(\sigma_1 \ge \sigma_2 \ge \dots \ge 0\) ,分别乘到 \( V^Tx \) 的各个维度上,由此可以写出不等式:
\[ \|Wx\| = \|\sigma_1 x_1+\sigma_2 x_2 + \dots\| \le \sigma_{\max} \|x\|, \]
其中 \( \sigma_{max} \) 是最大的奇异值。因此,Wx 变换后范数所能达到的最大缩放是由 \(\sigma_{\max}\) 决定的。
也就是说,我们可以通过控制矩阵 W 的最大奇异值来控制稳定性。
Wx 对方差的影响
方差反映的是各元素围绕均值的散布幅度。如果方差很大,意味着激活值的分布被摊得很宽,某些分量可能极端偏离均值过大导致不稳定;如果方差太小,所有神经元的输出近乎相同,网络也就失去了表达能力。所以和范数类似,我们期望输入输出的方差也稳定在一个适中的值附近,比如标准的 1。
相比于 L2 范数,方差考虑了向量的元素个数,假设 x 有 \( d_1 \) 个元素,而 W 则是 \( d_2 \times d_1 \) 形状
Wx 的某个元素 \( y_i \) 是 W 中对应行 \( W_i^T \) 与 x 的内积: \(y_i = W_i^T x\)
假设权重与输入相互独立,且 x 的均值为零、方差为 \(\sigma_x^2\),而 W 所有元素的均值为 0 方差为 \( \sigma_w \)
先求输出的期望: \(E[y_i] = E[\sum_{j=1}^{d_1} W_{ij} x_j]\) ,根据期望线性性质和 \( W_{ij} \) 与 \( x_j \) 的独立行,于是有 \( E[y_i]=0 \) 。
这个性质使得 \(y_i\) 的方差就等于它的平方期望: \( \text{Var}(y_i) = \mathbb{E}[y_i^2]+\mathbb{E}[y_i]^2 = \mathbb{E}[y_i^2] \)
而
\[ Var(y_i) = \mathbb{E}\left[\left(\sum_{j=1}^{d_1} W_{ij} x_j\right)^2\right] \]
把平方展开会出现平方项 \( W_{ij}^2 x_j^2 \) 和交叉项 \(W_{ij} W_{ik} x_j x_k\)(当 j ≠ k 时)。
但交叉项里各个变量都是独立的,于是期望可以代入,结果为 0 ;只剩下平方项的贡献:
\[ \text{Var}(y_i) = \sum_{j=1}^{d_1} \mathbb{E}[W_{ij}^2] \cdot \mathbb{E}[x_j^2] = \sum_{j=1}^{d_1} \sigma_w^2 \cdot \sigma_x^2 = d_1 \cdot \sigma_w^2 \cdot \sigma_x^2 \]
因为 W 的行是相互独立的同分布的随机变量(实际每个元素都是独立同分布,每个值独立地从某个分布里采样),所以这个就是 y 中各向量元素的方差。
以上表明,Wx 输出方差是输入方差的 \(n \sigma_w^2\) 倍。为了让信号在深层网络中稳定传播,如果期望输出方差依然保持 \( \sigma^2_x \),那么放大系数要刚好为 1,即满足 \(d_1 \sigma_w^2 = 1\),W 的方差需要被初始化为 \(\sigma_w^2 = 1/d_1\)。
这是 Xavier 初始化的核心思想,各类权重初始化方法会在后文有更多梳理,这里想说明的是,可以通过随机化方式来影响矩阵 W 对 x 方差的变换能力。
但前文说到,权重在梯度下降过程中不断变化的,初始化只能确保训练时最初的稳定性。
我们还要保证不断迭代后的权重仍然能够使得方差保持稳定,但这个过程中我们更难控制 \( W_{t+1} \) 的方差,后文会看到,在这种叠加修改 W 的场景中,能更加控制的是 W 的奇异值,于是就要想,是否可以用奇异值来影响方差呢?
这是要引入 RMS 范数。
RMS 范数
直接用 L2 范数还存在一个小问题:比如 W 是 \( \mathbb{R}^{d_2 \times d_1} \) 形状,如果我们要求范数在经过 Wx 变换后仍然保持稳定,比如期望 x 和 y 的范数都是 1:
- 在 d2 远小于 d1 的情况下,y 中的每个元素值就会非常小,极端情况甚至有些值会因为浮点误差而变成 0 。
- 反过来,在 d2 远大于 d1 的情况下, y 中的每个元素值就会非常大,极端情况出现 NaN 。
所以 L2 范数并不是一个好的控制稳定性的指标,但介绍它是因为它可以通过 W 的奇异值去影响。
而我们核心目标是控制方差(因为方差限制的是每个元素值,和元素个数无关),但目前我们只能在 W 随机初始化时通过 W 的方差控制输出向量方差。
这里我们遇到了一个 gap, 即想要用 W 奇异值控制范数从而控制输出向量方差,但却只能用 W 的方差去控制。
解决方法非常简单:只要对 L2 范数“归一化”得到 RMS 范数(root mean square norm) \( \frac{1}{\sqrt{d_1}}\| x\|_2 \) ,在均值为 0 时它的期望就是标准差(方差开根号),而它本身仅仅是 L2 范数除以一个常数,所以能用奇异值控制 L2 范数,就能用奇异值控制 RMS 范数(也就是标准差)。
而控制方差在一个适当值就等价于控制标准差在某个适当值,它们是等价的。
实际中我们对 RMS norm 的平方的期望保持为特定常数即可,而它就是方差:
以输出的 squared RMS norm 的期望为例,数学上写成 \( E[ \frac{1}{d_2} \sum y_i^2 ] \) = \( \frac{1}{d_2} \sum_{i=1}^{d_2} \mathbb{E}[y_i^2] \)
由于每个 \(y_i\) 的均值为 0 ,所以 \( E[y_i^2]=Var[y_i] \), 而每个方差都相同,于是结果就是上节计算过的方差 \( Var[y_i] = d_1 \sigma_w^2 \sigma_x^2 \)
它使得我们能够通过控制矩阵 W 的最大奇异值来控制输出向量的 L2 norm 从而控制乘以 \( \frac{1}{\sqrt{d_2}} \) 之后的 L2 norm ,也就是方差。
前两节中对 L2 norm 的影响公式为:
\[ \|Wx\| = \|\sigma_1 x_1+\sigma_2 x_2 + \dots\| \le \sigma_{\max} \|x\|, \]
替换成 RMS norm (的期望,我们只能保证期望而不能控制每个向量具体的 norm):
\[ \|Wx\|_{RMS} = \frac{1}{\sqrt{d_2}}\|\sigma_1 x_1+\sigma_2 x_2 + \dots\| \le \frac{1}{\sqrt{d_2}}\sigma_{\max} \|x\| \leq \frac{\sigma_{\max}}{\sqrt{d_2}} \sqrt{d_1} \|x\|_{RMS} \]
于是要保证 RMS norm 的期望稳定(标准差稳定), \( \sigma_{\max} \) 就应该尽量控制在 \( \sqrt{\frac{d_2}{d_1}} \) 。
现在遗留的问题就是我们如何在梯度更新的时候保持 W 的最大奇异值在这个稳定值附近。
如何控制 Ax 中各部分
随机初始化
前文已经提到过随机初始化对方差的影响,并且给出了 Xavier 的基本思路。这里先回顾并在加入了 ReLU 影响后推导出 Kaiming 初始化方法。
从行内积的方差传递出发,假设权重与输入独立,且输入 x 的各分量均值为零、方差为 \(\sigma_x^2 = 1\),权重 W 的各元素也独立同分布、均值为零、方差为 \(\sigma_w^2\)。在没有非线性激活函数的情况下,对于线性层 \(y = W x\),每个输出分量 \(y_i\) 的方差为
\[ \operatorname{Var}(y_i) = n \cdot \sigma_w^2 \cdot \sigma_x^2 = n \,\sigma_w^2 . \]
为了让输出方差维持在 1,就需要 \(n\,\sigma_w^2 = 1\),即 \(\sigma_w^2 = 1/n\)。这里的 n 是输入维度,通常记作 fan_in。如果只考虑反向传播的对称性,用输出维度 m 代替 n,会得到 \(\sigma_w^2 = 1/m\)。Xavier 初始化折中了前向与反向的需求,将权重方差设定为
\[ \sigma_w^2 = \frac{2}{n + m}, \]
在均匀分布或正态分布下按此方差进行采样。
实践中,每个线性层输出的分布会被激活函数修改,而最常使用的 ReLU 会显著改变信号的方差。
假设线性层的输出 y 已经通过合适的初始化使其各分量满足均值为零、方差为 \(\sigma_y^2\)。经过 ReLU 后,激活值 \(z = \max(0, y)\)。若 y 的分布关于零点对称(比如均值为零的正态分布),则负半轴的值被全部置零,正向部分保留。这带来的直接后果是,z 的方差大约减半。更严格地,若 \(y \sim \mathcal{N}(0, \sigma_y^2)\),则
\[ \operatorname{Var}(z) = \mathbb{E}[z^2] - (\mathbb{E}[z])^2 = \frac{1}{2} \sigma_y^2 - \left( \frac{1}{\sqrt{2\pi}} \sigma_y \right)^2 = \frac{\sigma_y^2}{2} \left(1 - \frac{1}{\pi}\right) \approx 0.34\,\sigma_y^2, \]
但为了保持初始化阶段的简洁有效,通常用“方差减半”来近似,即 \(\operatorname{Var}(z) \approx \frac{1}{2} \operatorname{Var}(y)\)。
现在我们想要的是经过 ReLU 后,z 的方差仍然维持与输入 x 相同的水平,比如 1。这要求进入 ReLU 之前的 y 的方差约为 2。回顾线性层的方差公式 \(\operatorname{Var}(y_i) = n \,\sigma_w^2 \cdot 1\),令其等于 2,立即得到
\[ \sigma_w^2 = \frac{2}{n}. \]
这就是 Kaiming 初始化(又称 He 初始化)的核心依据。相比 Xavier 初始化,它多乘了一个因子 2 来补偿 ReLU 造成的方差减半。与之对应,在实际实施时,如果采用正态分布,则标准差设为 \(\sqrt{2/n}\);若采用均匀分布,区间半径则按 \(\sqrt{6 \cdot 2/n} = \sqrt{12/n}\) 来限定。一
般来说,n 取输入维度 fan_in ,因为 ReLU 更影响前向传播(反向传播往往有残差层补偿),优先保障前向信号的稳定。对于带有负斜率的变体如 Leaky ReLU(负部分斜率为 a),补偿因子会调整为 \(2/(1+a^2)\),原理完全相同。
约束梯度下降中权重的方差
好的初始化至少可以使得在训练的前期整个网络数值保持稳定,但随着梯度不断被更新,它的影响力会越来越弱。
这时候根据梯度更新公式 \( \theta_{t+1} = \theta_t - \eta g(\nabla_\theta L) \) ,梯度(或关于梯度的函数)对参数不断累加在影响着网络的稳定性。
我们的目的是,给定 \( \theta \) ,它能够使得对正向 x 和反向梯度 g 在传播中保持稳定,那么我们希望 \( \theta+\Delta \theta \) 仍然能维持这个性质。
对损失进行最小化过程中,我们的目的不是一定要去遵循梯度的反方向走一小步,沿着梯度仅仅是因为它是当前点增长最大的方向,但这也只限于无穷小的范围内,给定一个固定步长 \( \eta \) ,并不是朝着梯度反方向走这一步就能达到最大的下降值,甚至可能在梯度负方向函数急剧反转,一步之后反而值变大了。
所以对于每一步优化,我们不单能对学习率进行选择,还能够对更新的方向进行取舍。
也就是说,单步更新中本身有一个足够大的搜索空间,在该空间里我们还能提出更多的约束,比如使得 \( \theta+\Delta \theta \) 仍然具有保持 RMS 范数恒定的性质。
于是单步梯度更新的目标包括:
尽可能使得目标函数有更大的降幅(由于变化量是负的,因此是最小化): \( \arg\min \Delta L(\theta) \)
找出 \( \Delta L(\theta) \) 最准确的方法是差分即 \( L(\theta + \Delta \theta) - L(\theta) \) ,但这意味着我们需要遍历(或者采样) \( \theta \) 的邻域,每次都计算出新的 \( L(\theta + \Delta \theta) \) 然后计算变化量,这对于深度学习是不切实际的,于是我们只能转向某种近似,而最简单高效的是一阶近似,它等于 \( g^T \Delta \theta \) 这里 \( g \) 是 \( \theta \) 的梯度向量(我们这里把所有参数看作列向量)
尽可能使得 \( \theta + \Delta \theta \) 仍然维持对数据转换的稳定性,这取决于我们希望保证哪种稳定性,比如如果 \( \theta \) 是矩阵,且保证它对输入输出数据方差的变换维持稳定,那么应该保持矩阵数值自身的方差尽可能不变,这样它才能延续着初始化时参数的性质。而即便 \( \theta \) 和 \( \Delta \theta \) 是独立的,要使得二者相加后的矩阵内数值的方差不变,意味着 \( \Delta \theta \) 的方差应该尽可能是 0 ,而且最好能保证均值也尽可能为 0, 因为初始化时参数均值为 0, 并且 relu 等常见的激活函数非线性区在 0 附近,因此均值为 0 能使得网络释放更强的非线性表示能力。
于是综合起来的优化目标是 \( \arg\min g^T \Delta \theta, \quad \arg\min Var(\Delta \theta) \quad s.t. E[\Delta \theta]=0 \)
但这里同时最大化降幅又最小化方差,二者不可能同时达到,因为最小化方差后 \( \Delta \theta \) 是个所有值都相同的矩阵,但它期望又是 0, 因此就是矩阵,此时降幅为 0 。
更合理的描述是每次允许方差小范围线性波动,这样从复杂度的角度来看,迭代 N 次后,方差变化量是 O(N), 控制在线性程度内。
即 \[ \arg\min g^T \Delta \theta, \quad s.t. Var(\Theta \theta)=\delta, E[\Delta \theta] = 0 \]
注意这里 \( \delta \) 是个超参,它本身应该设定地足够小我们才能用 \( g^T \Delta \theta \) 来近似变化量。
这里同时要满足期望和方差性质,还是不容易同时满足,于是我们转换问题,如果样本的方差在均值为 0 ,方差就是 \( \frac{1}{n} \sum x_i^2 \) ,这就是 RMS 范数的平方,于是更方便进行优化的写法是:
\[ \arg\min g^T \Delta \theta, \quad s.t. \frac{1}{\sqrt{n}} \| \Delta \theta \| = \delta \]
因为 \( g^T\Delta \theta \) 是关于 \( \Delta \theta \) 的线性函数,而范数又是凸函数,因此这是典型的一类优化问题,有现成的解法。
在求解这个问题前,我们先看如果是对扰动矩阵进行其他范数的约束,会得到哪些优化算法。
这里我们限定的是 \( \Delta \theta \) 的范数,不能它限制为 0, 而是让它不随着宽度变化而变化,因此我们需要的是 O(1) 复杂度
GD: 约束梯度下降中权重的 L2 范数
可以把对 \( \Delta \theta \) 的约束一般化:
\[ \Delta\theta = \arg\min_{\Delta\theta} \left[ g^T \Delta\theta + \frac{1}{\eta} d(\Delta\theta) \right] \] \( \eta > 0 \) 控制惩罚强度(\( \eta \) 越大,允许的步长越大),而 \( d(\cdot) \) 是衡量更新幅度的距离函数。
如果距离函数是一般的 L2 范数的平方,也就是欧几里得距离(平方),那么
\[ \arg\min g^T \Delta \theta, \quad s.t. \| \Delta \theta \|^2 = \delta \]
这种情况 RMS norm 也会被限制,因此这是一种间接手段。
可以改写成无条件约束的最优化形式: \[ \Delta\theta = \arg\min_{\Delta\theta} \left[ g^T \Delta\theta + \frac{1}{2\eta} \|\Delta\theta\|_2^2 \right] \]
解析解:对 \( \Delta\theta \) 求导并令其为零: \[ g + \frac{2}{2\eta} \Delta\theta = 0 \quad\Longrightarrow\quad \boxed{\Delta\theta = -\eta g} \]
这个结果告诉我们:最优更新方向严格沿着负梯度方向,步长与梯度大小成正比: \[ \theta_{t+1} = \theta_t - \eta g \]
这就是标准梯度下降(Gradient Descent)算法。
而如果是对 L2 范数约束(没有平方),经计算会得到:
\[ \theta_{t+1} = \theta_t - \eta \frac{g}{\|g\|^2} \]
这是归一化后的梯度下降。
SignGD: 约束梯度下降中权重的无穷范数
我们还可以通过限制向量中绝对值最大的那个值的范围来间接限制 \( \Delta \theta \) 的 RMS nrom 。
这也被称为无穷范数: 转成无约束问题: \[ \Delta\theta = \arg\min_{\Delta\theta} \left[ g^T \Delta\theta + \frac{1}{\eta} \|\Delta\theta\|_\infty \right] \]
要最小化需要对 \( \Delta \theta \) 求导,但向量最大绝对值的对向量的导数是多少?
从代数上这并不好计算。
更自然的思路是不转成无约束问题,直接考虑,在 \( \Delta \theta \) 的所有值都被限制在 \( [-\eta, \eta] \) 区间里时, \( g^T \Delta \theta \) 的最小值会取哪里?
因为 \( \Delta \theta \) 每个值都可正可负,而且是对称的,因此只要每个分量取和 g 中对应分量异号且绝对值最大,那么 \( g^T \Delta q \) 就会是一个负的很小的值。
于是可以得出最优的 \( \Delta \theta \) 就是:
\[ \Delta\theta = -\eta \cdot \text{sign}(g) \]
这个结果对应着符号梯度下降(Sign Descent / SignSGD)算法: \[ \theta_{t+1} = \theta_t - \eta \cdot \text{sign}(g) \]
注意这里更新方向并不是梯度的反方向,而只是负梯度向量所在象限里的一个向量,它也指示函数值减小的方向,只不过不是特定点最陡峭的方向。
RMSprop: 移动平均后的归一化
如果我们用梯度的平方进行归一化而不是当前梯度的平方根,就得到了 RMSprop
\[\begin{cases} r = \rho r + (1 - \rho) (\frac{\partial \mathcal{L} }{\partial W })^2 \\ w = w - \eta \frac{1}{\sqrt{r} + \varepsilon} \frac{\partial \mathcal{L} }{\partial W } \end{cases} \]
如果 \( \rho = 0 \) 即不考虑历史梯度的平方信息,就退化成了 signGD
Adam: RMSprop + Momentum
在 RMSprop 的基础上引入动量(梯度的移动平均),就得到了 Adam 优化器的雏形:
\[ \begin{cases} m = \beta_1 m + (1-\beta_1) \frac{\partial \mathcal{L}}{\partial W}, \\ v = \beta_2 v + (1-\beta_2) \left(\frac{\partial \mathcal{L}}{\partial W}\right)^2, \\ W = W - \eta \frac{m}{\sqrt{v} + \varepsilon} \end{cases} \]
一般超参取值为 \( \eta = 0.001, \beta_1 = 0.9, \beta_2 = 0.999 \) ;且 m 和 v 在初始时刻都设置为零, β1 和 β2 非常接近 1 会导致在最初的几步迭代中 m 和 v 会设置地非常小(但它应该就是接近第一次计算的梯度)。
为此,Adam 引入了偏置校正项,即用 1 减去 β 的 t 次方来除回去。于是,更严谨的形式应当是:
\begin{cases} m = \beta_1 m + (1-\beta_1) \frac{\partial \mathcal{L}}{\partial W}, \\ v = \beta_2 v + (1-\beta_2) \left(\frac{\partial \mathcal{L}}{\partial W}\right)^2, \\ \hat{m} = \frac{m}{1 - \beta_1^t}, \quad \hat{v} = \frac{v}{1 - \beta_2^t}, \\ W = W - \eta \frac{\hat{m}}{\sqrt{\hat{v}} + \varepsilon} \end{cases}这里 t 从 1 开始计数,在迭代第一步 m 和 v 就等于第一次计算到的梯度以及梯度平方;
而随着 t 增大,β^t 迅速趋近于零,校正项逐渐退化为 1,算法也就平滑地过渡到了未校正的形式。
从这里可以看出,一阶动量是用来提高一般收敛速度的,而二阶梯度是保证数值稳定性的。
AdamW: Adam+weight decay
AdamW 是 Adam 加上 weight decay :
\[ \theta_{t} = (1-\lambda) \theta_{t-1} - \eta \frac{m_{t}}{\sqrt{v_{t}} + \epsilon} \]
注意这不同于把 L2 正则项也考虑进来求梯度,因为那样的话 momentum 里也会有 L2 正则的梯度混进来。 这里我们要的就是权重衰减而不是一个加了正则项的新的损失函数。
embedding 和 de-embedding 层
embedding 层为了使得每个 token 的特征向量 RMS norm 期望为 1, 只需要 embedding 的每一列 RMSnorm 期望为 1, 于是初始化的时候可以设置为均值为 0; 方差为 1
最高层解码进入 softmax 层之前,为了保证 softmax 里数值稳定性,y=Wx 如果输入 x 的 RMSnorm 为 1, W 是一个很高(比如 5w 个 token 就有 5w 行,输出特征维度 1000, W 就是 5wx1000),要控制 W 的奇异值在 \( sqrt{\frac{5w}{1000}} = \sqrt{500} \) 内
从 row 角度,这保证 W 的每行(1000 维度)向量 RMSnorm 为 \( 1/\sqrt{1000} \) ,这样和 x 内积后单个值输出期望为 1 这种 row 视角其实是 F norm 视角,而奇异值视角是谱范数视角。 但对于这种高瘦的随机初始化矩阵,奇异值很接近,因此 F norm (所有奇异值的和)基本是谱范数(最大奇异值)的倍数,因此是等价的