勾配消失と勾配爆発

PyTorchによる中級ディープラーニング

Michal Oleszak

Machine Learning Engineer

勾配消失

  • 逆伝播中に勾配が徐々に小さくなる
  • 前方の層でパラメータ更新が小さくなる
  • モデルが学習できなくなる

勾配のサイズと層インデックスのグラフ:前方の層ほど勾配が小さい

PyTorchによる中級ディープラーニング

勾配爆発

  • 勾配が徐々に大きくなる
  • パラメータの更新量が大きすぎる
  • 学習が発散する

勾配のサイズと層インデックスのグラフ:前方の層ほど勾配が大きい

PyTorchによる中級ディープラーニング

不安定な勾配の解決策

  1. 適切な重みの初期化
  2. 適切な活性化関数
  3. バッチ正規化

 

 

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関数のプロット。ゼロ以下では水平、ゼロ以上では正の傾きを持つ。

  • デフォルトの活性化関数としてよく使われる
  • nn.functional.relu()
  • 負の入力でゼロ — ニューロンの死問題

ELU関数のプロット。ReLUに似ているが、負の領域から正の領域への遷移が滑らか。

  • nn.functional.elu()
  • 負の値でも非ゼロ勾配 — ニューロンの死を防ぐ
  • 平均出力がゼロ付近 — 勾配消失を緩和
PyTorchによる中級ディープラーニング

バッチ正規化

層の後に:

  1. 層の出力を正規化:

    • 平均を引く
    • 標準偏差で割る
  2. 学習パラメータで正規化出力をスケール・シフト

各層に最適な入力分布をモデルが学習:

  • 損失の減少が速くなる
  • 不安定な勾配を緩和
PyTorchによる中級ディープラーニング

バッチ正規化

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による中級ディープラーニング

練習しましょう!

PyTorchによる中級ディープラーニング

Preparing Video For Download...