梯度消失与爆炸

PyTorch 深度学习进阶

Michal Oleszak

Machine Learning Engineer

梯度消失

  • 反向传播中梯度逐层变小
  • 早期层参数更新很小
  • 模型难以学习

展示梯度大小与层索引关系的图:越早的层梯度越小

PyTorch 深度学习进阶

梯度爆炸

  • 梯度越来越大
  • 参数更新过大
  • 训练发散

展示梯度大小与层索引关系的图:越早的层梯度越大

PyTorch 深度学习进阶

解决不稳定梯度

  1. 合理的权重初始化
  2. 合适的激活函数
  3. 批量归一化

 

 

三步

PyTorch 深度学习进阶

权重初始化

layer = nn.Linear(8, 1)
print(layer.weight)
Parameter containing:
tensor([[-0.0195,  0.0992,  0.0391,  0.0212,
         -0.3386, -0.1892, -0.3170,  0.2148]])
PyTorch 深度学习进阶

权重初始化

良好初始化可确保:

  • 层输入方差 = 层输出方差
  • 梯度方差在层前后保持一致

 

实现方式取决于激活函数:

  • 对 ReLU 等,可用 He/Kaiming 初始化
PyTorch 深度学习进阶

权重初始化

import torch.nn.init as init

init.kaiming_uniform_(layer.weight)
print(layer.weight)
Parameter containing:
tensor([[-0.3063, -0.2410,  0.0588,  0.2664,
          0.0502, -0.0136,  0.2274,  0.0901]])
PyTorch 深度学习进阶

He / Kaiming 初始化

init.kaiming_uniform_(self.fc1.weight)
init.kaiming_uniform_(self.fc2.weight)
init.kaiming_uniform_(
  self.fc3.weight,
  nonlinearity="sigmoid",
)
PyTorch 深度学习进阶

He / Kaiming 初始化

import torch.nn as nn
import torch.nn.init as init

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(9, 16)
        self.fc2 = nn.Linear(16, 8)
        self.fc3 = nn.Linear(8, 1)


init.kaiming_uniform_(self.fc1.weight) init.kaiming_uniform_(self.fc2.weight) init.kaiming_uniform_( self.fc3.weight, nonlinearity="sigmoid", )




    def forward(self, x):
        x = nn.functional.relu(self.fc1(x))
        x = nn.functional.relu(self.fc2(x))
        x = nn.functional.sigmoid(self.fc3(x))
        return x







PyTorch 深度学习进阶

激活函数

ReLU 函数图:小于 0 时为 0,超过 0 后为正斜率直线。

  • 常用作默认激活
  • nn.functional.relu()
  • 负输入为 0 —— 神经元死亡问题

ELU 函数图:类似 ReLU,但在负区平滑过渡到正区。

  • nn.functional.elu()
  • 负值处梯度非零——缓解神经元死亡
  • 输出均值近 0——缓解梯度消失
PyTorch 深度学习进阶

批量归一化(BatchNorm)

在每层之后:

  1. 规范化该层输出:

    • 减去均值
    • 除以标准差
  2. 用可学习参数缩放和平移规范化结果

模型学习各层的最优输入分布:

  • 更快降低损失
  • 缓解不稳定梯度
PyTorch 深度学习进阶

批量归一化(BatchNorm)

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(9, 16)
        self.bn1 = nn.BatchNorm1d(16)

        ...


def forward(self, x): x = self.fc1(x) x = self.bn1(x) x = nn.functional.elu(x) ...
PyTorch 深度学习进阶

Passons à la pratique !

PyTorch 深度学习进阶

Preparing Video For Download...