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

Функція втрат дискримінатора

Час визначити функцію втрат для дискримінатора. Нагадаємо, завдання дискримінатора — класифікувати зображення як справжні або згенеровані. Відповідно, дискримінатор отримує втрати, якщо класифікує виходи генератора як справжні (мітка 1) або справжні зображення як підробні (мітка 0).

Визначте функцію disc_loss(), яка обчислює втрати дискримінатора. Вона приймає п'ять аргументів:

  • gen — модель генератора
  • disc — модель дискримінатора
  • real — вибірка справжніх зображень із тренувальних даних
  • num_images — кількість зображень у батчі
  • z_dim — розмір вхідного випадкового шуму

Ця вправа є частиною курсу

Глибоке навчання для зображень із PyTorch

Переглянути курс

Інструкції до вправи

  • Використайте дискримінатор для класифікації зображень fake і запишіть передбачення в disc_pred_fake.
  • Обчисліть компонент втрат для підробних зображень, викликавши criterion на передбаченнях дискримінатора для підробних зображень і тензорі нулів того самого розміру.
  • Використайте дискримінатор для класифікації зображень real і запишіть передбачення в disc_pred_real.
  • Обчисліть компонент втрат для справжніх зображень, викликавши criterion на передбаченнях дискримінатора для справжніх зображень і тензорі одиниць того самого розміру.

Інтерактивна практична вправа

Спробуйте виконати цю вправу, доповнивши цей зразок коду.

def disc_loss(gen, disc, real, num_images, z_dim):
    criterion = nn.BCEWithLogitsLoss()
    noise = torch.randn(num_images, z_dim)
    fake = gen(noise)
    # Get discriminator's predictions for fake images
    disc_pred_fake = ____
    # Calculate the fake loss component
    fake_loss = ____
    # Get discriminator's predictions for real images
    disc_pred_real = ____
    # Calculate the real loss component
    real_loss = ____
    disc_loss = (real_loss + fake_loss) / 2
    return disc_loss
Редагувати та запускати код