Функция потерь дискриминатора
Пришло время определить функцию потерь для дискриминатора. Вспомните, что задача дискриминатора — классифицировать изображения как настоящие или сгенерированные. Дискриминатор несёт потери, если он принимает выходные данные генератора за настоящие (метка 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