Kom igångKom igång gratis

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, generatormodellen
  • disc, diskriminatormodellen
  • real, ett urval av äkta bilder från träningsdata
  • num_images, antalet bilder i batchen
  • z_dim, storleken på det slumpmässiga brusets inmatning

Den här övningen är en del av kursen

Djupinlärning för bilder med PyTorch

Visa kurs

Övningsinstruktioner

  • Använd diskriminatorn för att klassificera fake-bilder och tilldela förutsägelserna till disc_pred_fake.
  • Beräkna förlustkomponenten för falska bilder genom att anropa criterion med 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 till disc_pred_real.
  • Beräkna förlustkomponenten för äkta bilder genom att anropa criterion med 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
Redigera och kör kod