Fonction de perte du discriminateur
Il est temps de définir la fonction de perte du discriminateur. Rappelez-vous que le rôle du discriminateur est de classer les images comme réelles ou fabriquées. Par conséquent, le générateur subit une perte s'il classe les sorties du générateur comme réelles (étiquette 1) ou les images réelles comme fabriquées (étiquette 0).
Définissez la fonction disc_loss() qui calcule la perte du discriminateur. Elle prend cinq paramètres :
gen, le modèle de générateurdisc, le modèle de discriminateurreal, un échantillon d'images réelles provenant des données d'entraînementnum_images, le nombre d'images dans le lotz_dim, la taille du bruit aléatoire en entrée
Cette activité fait partie du cours
Deep Learning pour les images avec PyTorch
Instructions de l’exercice
- Utilisez le discriminateur pour classer les images
fakeet affectez les prédictions àdisc_pred_fake. - Calculez la composante de perte pour les fausses images en appelant
criterionsur les prédictions du discriminateur pour les images fabriquées et sur un tenseur de zéros de même forme. - Utilisez le discriminateur pour classer les images
realet affectez les prédictions àdisc_pred_real. - Calculez la composante de perte pour les images réelles en appelant
criterionsur les prédictions du discriminateur pour les images réelles et sur un tenseur de uns de même forme.
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
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