НачатьНачать бесплатно

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

Наконец, вся работа по определению архитектур моделей и функций потерь приносит плоды: пришло время обучения! Ваша задача — реализовать и запустить цикл обучения GAN. Примечание: после первого батча данных добавлен оператор break, чтобы избежать долгого времени выполнения.

Два оптимизатора, disc_opt и gen_opt, инициализированы как оптимизаторы Adam(). Функции вычисления потерь, которые вы определили ранее, — gen_loss() и disc_loss() — уже доступны вам. Также подготовлен dataloader.

Напомним:

  • Аргументы disc_loss(): gen, disc, real, cur_batch_size, z_dim.
  • Аргументы gen_loss(): gen, disc, cur_batch_size, z_dim.

Это упражнение является частью курса

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

Посмотреть курс

Инструкции к упражнению

  • Вычислите потери дискриминатора с помощью disc_loss(), передав ей генератор, дискриминатор, выборку реальных изображений, текущий размер батча и размер шума 16 — именно в таком порядке. Результат присвойте переменной d_loss.
  • Вычислите градиенты на основе d_loss.
  • Вычислите потери генератора с помощью gen_loss(), передав ей генератор, дискриминатор, текущий размер батча и размер шума 16 — именно в таком порядке. Результат присвойте переменной g_loss.
  • Вычислите градиенты на основе g_loss.

Интерактивное практическое упражнение

Попробуйте выполнить это упражнение, дополнив этот пример кода.

for epoch in range(1):
    for real in dataloader:
        cur_batch_size = len(real)
        
        disc_opt.zero_grad()
        # Calculate discriminator loss
        d_loss = ____
        # Compute gradients
        ____
        disc_opt.step()

        gen_opt.zero_grad()
        # Calculate generator loss
        g_loss = ____
        # Compute generator gradients
        ____
        gen_opt.step()

        print(f"Generator loss: {g_loss}")
        print(f"Discriminator loss: {d_loss}")
        break
Редактировать и запускать код