Diskriminatorförlust
Nu är det dags att definiera förlusten för diskriminatorn. Kom ihåg att diskriminatorns uppgift är att klassificera bilder som äkta eller falska. Diskriminatorn drabbas alltså av en förlust om den klassificerar generatorns utdata som äkta (etikett 1) eller de äkta bilderna som falska (etikett 0).
Definiera funktionen disc_loss(), som beräknar diskriminatorförlusten. Den tar fem argument:
gen, generatormodellendisc, diskriminatormodellenreal, ett urval av äkta bilder från träningsdatanum_images, antalet bilder i batchenz_dim, storleken på det slumpmässiga brusets inmatning
Den här övningen är en del av kursen
Djupinlärning för bilder med PyTorch
Övningsinstruktioner
- Använd diskriminatorn för att klassificera
fake-bilder och tilldela förutsägelserna tilldisc_pred_fake. - Beräkna förlustkomponenten för falska bilder genom att anropa
criterionmed diskriminatorns förutsägelser för falska bilder och en tensor med nollor av samma form. - Använd diskriminatorn för att klassificera
real-bilder och tilldela förutsägelserna tilldisc_pred_real. - Beräkna förlustkomponenten för äkta bilder genom att anropa
criterionmed diskriminatorns förutsägelser för äkta bilder och en tensor med ettor av samma form.
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
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