【深度学习】数值稳定性和模型初始化

一、梯度爆炸 & 梯度消失

1:梯度爆炸(Exploding Gradients)

现象

loss 突然变成 nan或 loss 出现巨大跳变,数据变得非常大

原因:反向传播时梯度连乘,某一步梯度过大,更新后的权重直接溢出 / 越过最优解。

2:梯度消失(Vanishing Gradients)

现象

loss 几乎不变,一动不动 训练准确率永远卡在某个低值,怎么调 lr 都没用

原因 :反向传播连乘 ∂h/∂x 都远小于 1,经过几十层后就趋近于 0,前面层的权重根本接收不到有效梯度,前面的层等于废了

这两个其实是同一件事的两个极端:深层网络里,数值在连乘下要么越变越小,要么越变越大。


二、数学上看清:问题就出在"连乘"

一个 L 层的 MLP

每层做线性变换 + 激活(这里先用恒等激活 σ(x)=x 简化,只看线性部分):

z(l)=W(l)z(l−1)+b(l) \mathbf{z}^{(l)} = \mathbf{W}^{(l)} \mathbf{z}^{(l-1)} + \mathbf{b}^{(l)} z(l)=W(l)z(l−1)+b(l)

正向传播时,最终输出对输入的依赖,本质是所有权重矩阵的连乘

z(L)=W(L)W(L−1)⋯W(1)x \mathbf{z}^{(L)} = \mathbf{W}^{(L)}\mathbf{W}^{(L-1)} \cdots \mathbf{W}^{(1)} \mathbf{x} z(L)=W(L)W(L−1)⋯W(1)x

反向传播时梯度也是连乘

设线性激活下,第 ttt 层损失的梯度为:

∂L∂z(t)=W(t)W(t−1)⋯⏟一路连乘∂L∂z(L) \frac{\partial L}{\partial \mathbf{z}^{(t)}} = \underbrace{\mathbf{W}^{(t)}\mathbf{W}^{(t-1)} \cdots }_{\text{一路连乘}} \frac{\partial L}{\partial \mathbf{z}^{(L)}} ∂z(t)∂L=一路连乘 W(t)W(t−1)⋯∂z(L)∂L

权重矩阵连乘,导致梯度要么指数爆炸、要么指数消失。

简例

100 层,每层只有一个权重标量 www,所有层一样:

w100 w^{100} w100

  • 若 w=0.5w = 0.5w=0.5:0.5100≈7.9×10−310.5^{100} \approx 7.9\times10^{-31}0.5100≈7.9×10−31 → 直接消失
  • 若 w=1.5w = 1.5w=1.5:1.5100≈4×10171.5^{100} \approx 4\times10^{17}1.5100≈4×1017 → 直接爆炸
  • 只有当 www 恰好约为 1 时才稳定

结论很清晰:在"非 1 附近"的连乘下,深层网络必然数值失稳。 而网络初始化时,权重是随机给的------不保证在会导致稳定的范围

这就是为什么初始化如此重要:好的初始化让每一层的输出/梯度保持在合理尺度,从而让连乘不至于爆炸或消失。


三、激活函数也在"帮倒忙"

上面假设了恒等激活(σ(x)=x)。但真实激活会进一步放大问题,最典型的是 sigmoid / tanh

Sigmoid 的饱和区

Sigmoid 函数为 σ(x)=11+e−x\sigma(x) = \frac{1}{1+e^{-x}}σ(x)=1+e−x1,它的导数:

σ′(x)=σ(x)(1−σ(x)) \sigma'(x) = \sigma(x)(1-\sigma(x)) σ′(x)=σ(x)(1−σ(x))

它的取值最大只有 0.25,而且:

x σ(x) σ'(x)
0 0.5 0.25(最大)
5 ≈0.993 ≈0.007(很小)
-5 ≈0.007 ≈0.007(很小)

关键 :只要输入 xxx 稍微跑到 ±5 以外(饱和区),导数就趋近 0。反正的 0.25 甚至更小,每层梯度再乘一次小于 1 的数,深度稍微一深梯度必然消失。

这就是为什么现在用的 ReLU 成了隐藏层默认

ReLU(x)=max⁡(0,x),ReLU′(x)={1x>00x≤0 \text{ReLU}(x) = \max(0,x), \quad \text{ReLU}'(x) = \begin{cases}1 & x>0\\ 0 & x\le0\end{cases} ReLU(x)=max(0,x),ReLU′(x)={10x>0x≤0

  • 正半轴导数为 1 → 不缩小梯度 → 缓解梯度消失
  • 计算极快
  • 缺点:负半轴全为 0,可能导致神经元"死亡"(永远输出 0),但通常用很小的 lr 缓解

小结:选对激活函数(ReLU)本身,就是一种"数值稳定性"策略。


四、解法一:合理的参数初始化

现在进入正题。既然问题出在连乘尺度不对,那第一道防线就是把权重初始化的尺度搞对

4.1 错误示范:全 0 初始化

复制代码
如果所有 w = 0:
→ 所有神经元输出相同
→ 梯度相同
→ 所有神经元变得一模一样(对称性永远打破不了)
→ 网络退化成一个"笨重神经元",loss 不再下降

结论:权重绝不能全初始化为 0(偏置 b 可以)。

4.2 朴素随机初始化的问题

python 复制代码
net[t] 的权重 = 从标准正态 N(0,1) 随机取

问题:分布固定,但没考虑层的宽度(输入个数) 。层越宽,加权和 ∑iwixi\sum_i w_i x_i∑iwixi 的标准差越大,越容易爆。

4.3 让"方差"保持稳定

我们要让每层输出的方差尽量不变------这正是后面所有初始化方法的数学出发点。

设输入 xix_ixi 独立且方差为 v(x)v(x)v(x),权重 wi∼N(0,α2)w_i \sim N(0, \alpha^2)wi∼N(0,α2) 独立同分布。一个神经元的加权和:

z=∑i=1nwixi z = \sum_{i=1}^{n} w_i x_i z=i=1∑nwixi

方差为:

v(z)=∑i=1nv(wi)v(xi)=∑i=1nα2⋅v(x)=n α2 v(x) v(z) = \sum_{i=1}^{n} v(w_i) v(x_i) = \sum_{i=1}^{n} \alpha^2 \cdot v(x) = n\,\alpha^2\, v(x) v(z)=i=1∑nv(wi)v(xi)=i=1∑nα2⋅v(x)=nα2v(x)

想让 v(z)≈v(x)v(z) \approx v(x)v(z)≈v(x)(输出方差和输入方差一致),就需要

n α2=1  ⇒  α=1n n\,\alpha^2 = 1 \;\Rightarrow\; \alpha = \frac{1}{\sqrt{n}} nα2=1⇒α=n 1

也就是说,权重应初始化为 N(0,1n)N(0, \frac{1}{n})N(0,n1) 尺度 ,nnn 是本层的输入神经元个数。这就是著名的 Xavier(Glorot)初始化 的核心思想!

初始化要让"每层的方差保持稳定",就必须按层宽开根号来决定权重的标准差。


五、两大经典初始化

5.1 Xavier 初始化(适合 Sigmoid / Tanh / 线性激活)

同时考虑前向(输入个数 ninn_{in}nin)和反向(输出个数 noutn_{out}nout),取两者的平均:

α=2nin+nout \alpha = \sqrt{\frac{2}{n_{in} + n_{out}}} α=nin+nout2

权重从均匀分布 U(−6nin+nout,+6nin+nout)U(-\sqrt{\frac{6}{n_{in}+n_{out}}}, +\sqrt{\frac{6}{n_{in}+n_{out}}})U(−nin+nout6 ,+nin+nout6 ) 或正态 N(0,2nin+nout)N(0, \frac{2}{n_{in}+n_{out}})N(0,nin+nout2) 里采样。

PyTorch 一行:

python 复制代码
def init_weights(m):
    if type(m) == nn.Linear:
        nn.init.xavier_uniform_(m.weight)
net.apply(init_weights)

5.2 He / Kaiming 初始化(适合 ReLU 系)

Xavier 是在"激活前后方差相等"假设下推导的。但 ReLU 会把负半轴砍成 0,实际传递的方差缩小了一半 ,所以需要把系数加倍来补偿:

α=2nin(或2nin+nout) \alpha = \sqrt{\frac{2}{n_{in}}} \quad(\text{或} \sqrt{\frac{2}{n_{in}+n_{out}}}) α=nin2 (或nin+nout2 )

PyTorch 一行:

python 复制代码
def init_weights(m):
    if type(m) == nn.Linear:
        nn.init.kaiming_uniform_(m.weight, nonlinearity='relu')
net.apply(init_weights)

六、完整实验:亲眼看初始化有多重要

用 MLP 在 Fashion-MNIST 上,对比三种初始化,观察训练情况:

python 复制代码
import torch
import torch.nn as nn
import torchvision
from torch.utils.data import DataLoader

batch_size = 128
train_ds = torchvision.datasets.FashionMNIST(root='./data', train=True, download=True,
                                             transform=torchvision.transforms.ToTensor())
test_ds  = torchvision.datasets.FashionMNIST(root='./data', train=False, download=True,
                                             transform=torchvision.transforms.ToTensor())
train_iter = DataLoader(train_ds, batch_size, shuffle=True)
test_iter  = DataLoader(test_ds, batch_size)

loss = nn.CrossEntropyLoss()

def build_net():
    return nn.Sequential(
        nn.Flatten(),
        nn.Linear(784, 256), nn.ReLU(),
        nn.Linear(256, 256), nn.ReLU(),
        nn.Linear(256, 256), nn.ReLU(),
        nn.Linear(256, 10),
    )

def init_zero(m):
    if type(m) == nn.Linear: nn.init.zeros_(m.weight)
def init_xavier(m):
    if type(m) == nn.Linear: nn.init.xavier_uniform_(m.weight)
def init_kaiming(m):
    if type(m) == nn.Linear: nn.init.kaiming_uniform_(m.weight, nonlinearity='relu')

def run(init_fn, epochs=5):
    net = build_net()
    net.apply(init_fn)
    trainer = torch.optim.SGD(net.parameters(), lr=0.1)
    for epoch in range(epochs):
        net.train()
        total, correct = 0, 0
        for X, y in train_iter:
            trainer.zero_grad()
            l = loss(net(X), y)
            l.backward(); trainer.step()
            correct += (net(X).argmax(1) == y).sum().item()
            total += y.numel()
        # 测试准确率
        net.eval()
        t_c, t_t = 0, 0
        with torch.no_grad():
            for X, y in test_iter:
                t_c += (net(X).argmax(1) == y).sum().item()
                t_t += y.numel()
        print(f'epoch {epoch+1}: train {correct/total:.4f}, test {t_c/t_t:.4f}')

print('=== 全 0 初始化 ===')
run(init_zero)     # 预期:loss 几乎不动,准确率卡在 ~10%
print('=== Xavier 初始化 ===')
run(init_xavier)   # 预期:能训练,但可能略慢
print('=== He/Kaiming 初始化 ===')
run(init_kaiming)  # 预期:训练最顺畅,测试最好

七、还有其他"稳定数值"的手段(进阶)

初始化是最基础的一道防线,但光靠它还不够。工业界还有几个更强的手段:

手段 作用 感受
批量归一化(BatchNorm) 每层把输出拉回标准分布,强行稳住尺度 算是"用魔法打败魔法",最常用之一
残差连接(ResNet) 跳跃连接,让梯度有"高速公路"直达,缓解深层梯度消失 ResNet 靠它堆到 152 层
层归一化(LayerNorm) Transformer 用的归一化,对整层归一 NLP 标配
梯度裁剪 梯度过大时强行截断,防爆炸 训练 RNN 必备
相关推荐
冬奇Lab2 小时前
Code Agent 解剖(20):从零扩展——给 agent 加一个新工具
人工智能·开源
论文复现现场2 小时前
RTX 4090 24GB 能跑 Qwen3.8-27B 吗?单卡显存计算与云端部署指南
人工智能·python·云计算·llama·gpu算力
七牛云行业应用2 小时前
Harness Engineering 是什么:从“写提示词“到“设计 Agent 边界“的工程方法论
人工智能·agent·ai编程
科技拓维者2 小时前
短视频配音音效去哪里找比较好?短视频创作者音效指南
人工智能·音视频
cjy0001112 小时前
2026AI 产品如何从一次性 Demo 变成可重复使用的业务工具?
前端·人工智能·fde
大模型码小白2 小时前
AI大模型接入SDK:人工智能核心概念与发展史
大数据·运维·人工智能·安全·prompt
杀生丸学AI2 小时前
【动态重建】Flow4DGS-SLAM:基于光流引导的4DGS-SLAM算法
人工智能·三维重建·扩散模型·4dgs·动态重建
IT_陈寒2 小时前
React的状态更新坑得我差点加班到天亮
前端·人工智能·后端
腾视科技-AI2 小时前
腾视科技TS-NV-P200车载系列AI边缘算力盒子:引领车路协同新时代,赋能多元场景应用
人工智能·科技·ai·ai算力·ai边缘算力盒子·腾视科技·ai算力盒