Затухающие и взрывные градиенты

Глубокое обучение на 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. При значениях меньше нуля линия горизонтальна; при значениях выше нуля линия имеет положительный наклон.

  • Стандартная функция активации по умолчанию
  • 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...