ПочатиПочніть безкоштовно

Цикл навчання

Нарешті вся ваша робота над архітектурами моделей і функціями втрат дає результат: час навчати! Ваше завдання — реалізувати й запустити цикл навчання 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
Редагувати та запускати код