ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

多隐层网络梯度消失与Xavier、He初始化的数理推导

多隐层网络梯度消失与Xavier、He初始化的数理推导 多隐层网络的训练困难十次里有九次出在梯度在层与层之间传递时出了问题。最近帮一位朋友排查一个四层全连接网络结构不复杂数据也正常但损失就是卡在某个值附近下不去。我当时第一反应不是调学习率而是让他把每一层的梯度范数打印出来看一眼。果然靠近输出层的梯度量级在1e-2靠近输入层的梯度已经掉到1e-7。这种五个数量级的差距就是典型的梯度信号在多个隐层之间被“吃”掉了。这篇是“深度学习多隐层架构数理逻辑浅析”系列第十三篇的第6部分重点聊清楚一个事多隐层架构里的梯度传播是如何被矩阵乘积和激活函数导数共同影响的以及我们用Xavier、He初始化时背后的数理依据到底从哪来。内容适合那些已经能跑通模型但对“为什么这么调”还没有建立起直觉的读者。我会从复合函数和线性代数讲起逐步推到梯度、方差、初始化再落回到实际训练中的排查手段。1. 多隐层架构的数理建模先从复合函数说起1.1 层与层之间到底在算什么东西多隐层网络说白了就是把很多个函数套在一起。输入经过第一层输出变成中间表达再作为第二层的输入再经过一层直到最后输出。用符号写出来就是x₀ 是原始输入x₁ φ(W₁x₀ b₁)x₂ φ(W₂x₁ b₂)一直到 x_L φ(W_L x_{L-1} b_L)。这里的下标L代表总层数φ是激活函数。每一层做的事情就是一次仿射变换加一个非线性映射。仿射变换负责拉伸、旋转、缩放空间里的点激活函数负责“折叠”空间。两者配合才能把原本线性不可分的数据变成某种更容易被后续层处理的结构。数理逻辑上这个嵌套结构最大的价值在于表达能力的增长。一个只有一层隐层的网络理论上可以逼近任何连续函数但需要极大的宽度而多隐层网络用深度来换取更紧凑的参数表达。为什么不直接堆宽度因为参数数量会爆炸而且泛化性能通常不如同等参数量下的深网络。1.2 仿射变换与非线性折叠的线性代数解释从线性代数角度看一个全连接层就是矩阵乘法加偏置然后接非线性函数。如果去掉非线性函数不管堆多少层因为矩阵乘法满足结合律多层线性变换最终可以压缩成单个矩阵乘法。也就是说一个没有激活函数的深度网络表达能力等同于一个线性分类器这个问题在逻辑上就已经宣判了深度架构的死刑。所以激活函数不是“加一点非线性”这么简单而是深度网络能够处理非线性问题的基石。不同的激活函数决定了“折叠”的方式。比如tanh在0附近导数接近1会把空间光滑地压缩到[-1,1]ReLU则是把负半轴全部压到0形成一种单侧的“裁剪”。这些看似细节的选择会直接影响后面梯度传播的量级。也可以用揉面团的类比来理解。线性层就像把面团拉伸、压扁激活函数就像把面团折叠起来。如果只拉伸不折叠面团永远是一张薄饼只有反复折叠才能形成层次分明、结构复杂的起酥。深度网络每一层都在做这样一次“拉伸折叠”而数理逻辑要回答的就是这个过程在梯度回传时会不会把信号弄丢。2. 梯度从后往前传链式法则背后的矩阵结构2.1 反向传播的雅可比乘积长什么样训练深度网络的默认算法是反向传播本质就是求损失函数对每个参数的偏导时不断套用链式法则。对第l层的权矩阵W_l梯度可以写成∂L / ∂W_l (∂L / ∂z_l) · h_{l-1}ᵗ而误差信号 δ_l ∂L / ∂z_l 不能凭空得到它必须从后一层传回来。具体公式是δ_l (W_{l1}ᵗ · δ_{l1}) ⊙ φ(z_l)这里的 ⊙ 表示逐元素乘积φ(z_l) 是激活函数在z_l处的导数向量。如果我把这个过程展开到更早的层会看到梯度信号实际上是很多个矩阵和向量连续相乘的结果。更准确地说误差从第L层传到第1层需要乘上δ₁ D₁ W₂ᵗ D₂ W₃ᵗ ... D_{L-1} W_Lᵗ δ_L其中D_l是对角矩阵对角线是φ(z_l)的各个分量。这是多隐层架构里最核心的数学结构。你可以把它理解为梯度信号要穿过一道又一道的门每道门由权矩阵的转置和激活函数导数共同构成。任何一个门的“通量”太小信号就会被削弱任何一个门“通量”太大信号就会被放大。2.2 为什么乘积矩阵范数决定了梯度的生死把上面的乘积看成一个整体矩阵网络能否稳定训练很大程度上取决于这个整体矩阵的谱范数或奇异值分布。如果这个矩阵在某个方向上的最大奇异值小于1那么经过足够多层之后梯度在该方向上的分量会指数级收缩如果大于1就会指数级膨胀。这就是“梯度消失”和“梯度爆炸”的数学来源和具体损失函数无关只和网络结构、初始化参数以及激活函数的选择有关。这里要特别强调一个容易忽略的细节激活函数导数的缩放效应。sigmoid的导数最大只有0.25tanh的导数最大是1ReLU在正区间导数恒为1。如果激活函数使用的是sigmoid即使权重矩阵是正交矩阵每一层梯度也会至少乘以0.25。假设网络有10层光这一项就是0.25的10次方约等于9.5e-7梯度直接消失。这也是为什么现在很少见到深度网络用sigmoid作为隐层激活函数。所以说初始化策略本质上是在“开门”的时候把门的尺寸调到一个合适范围。不能太大否则梯度爆炸不能太小否则梯度消失。下一节就来拆解这个平衡是怎么用方差守恒算出来的。3. 初始化策略的数理推导方差守恒是怎么算出来的3.1 前向传播的方差约束为了让信息从输入层顺利传到输出层我们希望每一层输出的方差保持在同一量级。假设输入的每个特征独立同分布均值为0方差为σ²。第l层有n_l个神经元权矩阵W_l的元素独立同分布均值为0方差为σ_w²。先不考虑激活函数或认为初始化附近激活函数近似线性我们来算下一层神经元的输入方差。一个神经元的输入是上一层输出的加权和z Σ w_i x_i。由于w_i与x_i相互独立且均值为0这个加权和的方差就是每个项方差之和即 n_l · σ_w² · Var(x)。要让这一层输出的方差和上一层输入的方差保持近似相同就要求 n_l · σ_w² ≈ 1。所以理想情况下 σ_w 1 / sqrt(n_l)。这里的n_l是当前层的神经元数量也就是俗称的fan_in。刚才的推导忽略了一点如果激活函数是tanh它在0附近的导数约等于1输入输出方差基本不变。但如果激活函数是ReLU负半轴会把一半的神经元输出变成0方差就会缩小。所以ReLU要额外把方差放大。3.2 反向传播的方差约束与Xavier/He初始化的由来前向只是其中一半梯度还得从输出层往输入层传。反向传播时第l层回传梯度的方差也和权重方差的平方有关。用类似的推导可以得出要让反向梯度的方差不衰减需要 n_{l1} · σ_w² ≈ 1这里的n_{l1}是下一层的神经元数量也就是fan_out。问题来了前向需要1/sqrt(fan_in)反向需要1/sqrt(fan_out)。当一个层的输入维度和输出维度不相等时这两个约束不可能同时精确满足。于是Xavier初始化取了一个折中方案把权重方差设为2除以两者之和σ_w² 2 / (n_l n_{l1})这就是Glorot初始化也是tanh、sigmoid激活下最常用的初始化方式。对于ReLU由于负半轴导数为0信息有一半被截断He初始化把方差放大两倍σ_w² 2 / n_l也就是权重从均值为0、标准差为sqrt(2/n_l)的正态分布中采样。这样前向和反向都能获得大致稳定的梯度量级。这里我想强调一个容易犯的错很多人以为只要用了这些初始化公式网络就一定稳。其实初始化只是初始条件训练过程中权重会变化矩阵乘积的奇异值分布也会漂移。所以初始化解决的是“起跑线”问题而不是“全程护航”问题。它能把梯度量级控制住让训练一开始不至于发散或停滞。4. 数学结论落地到训练诊断与调参的实操路径4.1 如何用梯度范数判断网络是否健康理解了上面的推导诊断一个深度网络就不需要靠瞎猜了。我的习惯是在训练刚开始的几个step把每一层权重的梯度L2范数打出来。如果相邻层的梯度范数比例保持在一个不太离谱的范围比如0.5到2之间说明信号还能正常流动。如果比例长期是0.01甚至更小那基本就是梯度消失如果连续几层比例大于10那就要小心梯度爆炸。还可以做一个更细的检查把梯度范数按层画出来观察是不是从输出到输入呈现平滑的指数衰减或指数增长。如果是平滑的大概率是权重矩阵谱性质主导的可以通过调整初始化方差或改用残差连接来解决。如果某个特定层的梯度特别异常往往问题出在这一层的激活函数或前后层维度不匹配上。还有一个很有用的快速实验把激活函数换成线性激活也就是φ(x)x保持网络结构和数据不变先跑几十步。如果损失能明显下降说明数据和优化器没问题然后再把激活函数逐个加回去看是哪个位置开始出现梯度衰减。这个“二分法”能快速定位问题层比盲目调学习率高效得多。4.2 我踩过的几个初始化相关的坑第一个坑是bias初始化。很多人只初始化权重bias直接设成0这本身没有问题但如果网络某一层使用了ReLU所有神经元输入落在负区间的概率很大导致整个层梯度为0。因为ReLU负半轴导数为0一旦神经元输出为0累积误差无法回传神经元就“死”了。所以我一般建议ReLU网络的bias初始化为0.01这种很小的正值或者在后来的训练中发现大量dead unit时先把bias调大一点再检查学习率。第二个坑是忘记配合输入数据的分布。Xavier和He初始化都假设输入特征方差在1附近但如果你直接把原始像素值0到1的数据丢进去或者没有做标准化方差约束就失效了。我在某个实验里遇到过输入特征方差达到100结果初始激活值巨大梯度直接爆炸。解决方式很简单把输入标准化到零均值、单位方差或者按输入方差反过来缩小初始化权重。第三个坑是学习率和初始化不匹配。初始化方差略微偏大学习率又设置得偏高第一轮更新可能直接让loss变成NaN。这不能只怪学习率也可能怪初始化。我的经验是如果模型第一轮就NaN先检查初始化标准差再检查学习率。很多自适应优化器对初始化仍然敏感不要以为亚当什么都管。5. 再多一点数理直觉激活函数导数与残差连接的数学意义5.1 激活函数导数如何影响梯度缩放前面已经提到Sigmoid导数的最大值是0.25这是天然梯度衰减因子。我们再用更定量的方式看假设一个10层的网络每层权重矩阵的谱范数恰好为1如果激活函数是tanh在0附近导数也是1那么梯度乘积的范数理论上可以保持稳定。但如果激活函数是sigmoid即使每层矩阵是正交矩阵梯度也会乘上(0.25)^10这在数学上必然导致梯度消失。所以在设计一个深层网络时我的首要选择就是激活函数导数范围要包含1。ReLU在正区间导数为1负区间为0这带来一个问题不是梯度消失而是神经元死亡。Leaky ReLU在负区间给了一个小的斜率比如0.01等于把死亡概率降下来了导数都还在1附近徘徊。这些改动背后都是同一个数理逻辑尽可能让雅可比乘积的谱半径保持在1附近。5.2 残差连接为什么能缓解乘积效应残差连接在训练上百层的网络时几乎是标配。从数学上看残差块构造了一个新的映射y x F(x)。在反向传播时梯度除了要经过F这条路的雅可比乘积还有一条恒等路径直接把梯度传回去。也就是说即使F的雅可比矩阵的谱范数很小恒等路径仍然保留了完整梯度。这相当于在乘积结构中加入了一个“旁路”让梯度不必依赖连续矩阵乘积就能回传。如果将残差网络展开会发现它为梯度提供了多条不同长度的路径长短不一。所有路径加起来的平均梯度不再像纯连乘那样指数级衰减而是更接近于一个混合模型的平均值。这也是为什么残差网络在数学上比纯全连接网络更稳的原因。理解了这层我就不再觉得残差连接只是工程上的小技巧而是对雅可比乘积结构的一次系统性改造。最后说一点个人体会。看论文时很多人只记住结论比如“用He初始化”“加残差”却不关心这些建议背后的数理逻辑。但实际调模型时遇到的坑往往就出在推导被忽略的假设上。比如输入分布变了、激活函数换了原来的推荐就不一定适用。所以我现在每接触一个新的网络结构第一件事就是手动画一遍梯度传播路径看每个环节的雅可比乘积大概是多少。这个习惯救了我很多次也希望你试一下。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进