Обучение GAN

Глубокое обучение для работы с изображениями на PyTorch

Michal Oleszak

Machine Learning Engineer

Цель генератора

 

 

Схема работы GAN.

  • Цель: генерировать фейки, которые обманут дискриминатор
  • Идея: использовать дискриминатор для оценки качества генератора
  • Выход генератора классифицируется дискриминатором как:
    • Настоящий (метка 1) — хорошо, малые потери
    • Фейковый (метка 0) — плохо, большие потери
Глубокое обучение для работы с изображениями на PyTorch

Потери генератора

def gen_loss(gen, disc, num_images, z_dim):

noise = torch.randn(num_images, z_dim)
fake = gen(noise)
disc_pred = disc(fake)
criterion = nn.BCEWithLogitsLoss()
gen_loss = criterion( disc_pred, torch.ones_like(disc_pred) ) return gen_loss
  • Задать случайный шум
  • Сгенерировать фейковое изображение
  • Получить предсказание дискриминатора для фейка
  • Применить критерий бинарной кросс-энтропии (BCE)
  • Потери генератора: BCE между предсказаниями дискриминатора и тензором единиц
Глубокое обучение для работы с изображениями на PyTorch

Цель дискриминатора

 

 

Схема работы GAN.

  • Цель: правильно классифицировать фейки и настоящие изображения
  • Выходы генератора должны классифицироваться как фейк (метка 0)
  • Настоящие изображения должны классифицироваться как настоящие (метка 1)
Глубокое обучение для работы с изображениями на PyTorch

Потери дискриминатора

def disc_loss(gen, disc, real, num_images, z_dim):

criterion = nn.BCEWithLogitsLoss()
noise = torch.randn(num_images, z_dim)
fake = gen(noise)
disc_pred_fake = disc(fake)
fake_loss = criterion( disc_pred_fake, torch.zeros_like(disc_pred_fake) )
disc_pred_real = disc(real)
real_loss = criterion( disc_pred_real, torch.ones_like(disc_pred_real) )
disc_loss = (real_loss + fake_loss) / 2 return disc_loss
  • Задать критерий бинарной кросс-энтропии
  • Сгенерировать входной шум для генератора
  • Сгенерировать фейки
  • Получить предсказания дискриминатора для фейковых изображений
  • Вычислить компонент потерь для фейков
  • Получить предсказания дискриминатора для настоящих изображений
  • Вычислить компонент потерь для настоящих изображений
  • Итоговые потери — среднее между компонентами для настоящих и фейковых изображений
Глубокое обучение для работы с изображениями на PyTorch

Цикл обучения GAN

for epoch in range(num_epochs):
    for real in dataloader:
        cur_batch_size = len(real)


disc_opt.zero_grad()
disc_loss = disc_loss( gen, disc, real, cur_batch_size, z_dim=16)
disc_loss.backward() disc_opt.step()
gen_opt.zero_grad()
gen_loss = gen_loss( gen, disc, cur_batch_size, z_dim=16)
gen_loss.backward() gen_opt.step()
  • Итерация по эпохам и батчам реальных данных, вычисление размера текущего батча
  • Обнуление градиентов оптимизатора дискриминатора
  • Вычисление потерь дискриминатора
  • Вычисление градиентов дискриминатора и шаг оптимизации
  • Обнуление градиентов оптимизатора генератора
  • Вычисление потерь генератора
  • Вычисление градиентов генератора и шаг оптимизации
Глубокое обучение для работы с изображениями на PyTorch

Давайте потренируемся!

Глубокое обучение для работы с изображениями на PyTorch

Preparing Video For Download...