Kom igångKom igång gratis

Träningsloop

Nu när du har definierat modellarkitekturerna och förlustfunktionerna är det dags att sätta ihop allt: träningen kan börja! Din uppgift är att implementera och köra GAN:ets träningsloop. Observera att en break-sats är placerad efter den första datamängden för att undvika lång körtid.

De två optimerarna, disc_opt och gen_opt, har initierats som Adam()-optimerare. Funktionerna för att beräkna förlusterna som du definierade tidigare, gen_loss() och disc_loss(), finns tillgängliga för dig. En dataloader är också förberedd åt dig.

Kom ihåg att:

  • disc_loss() tar följande argument: gen, disc, real, cur_batch_size, z_dim.
  • gen_loss() tar följande argument: gen, disc, cur_batch_size, z_dim.

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

Djupinlärning för bilder med PyTorch

Visa kurs

Övningsinstruktioner

  • Beräkna diskriminatorns förlust med disc_loss() genom att skicka in generatorn, diskriminatorn, urvalet av riktiga bilder, aktuell batchstorlek och brusdimensionen 16, i den ordningen, och tilldela resultatet till d_loss.
  • Beräkna gradienter med hjälp av d_loss.
  • Beräkna generatorns förlust med gen_loss() genom att skicka in generatorn, diskriminatorn, aktuell batchstorlek och brusdimensionen 16, i den ordningen, och tilldela resultatet till g_loss.
  • Beräkna gradienter med hjälp av g_loss.

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

for epoch in range(1):
    for real in dataloader:
        cur_batch_size = len(real)
        
        disc_opt.zero_grad()
        # Calculate discriminator loss
        d_loss = ____
        # Compute gradients
        ____
        disc_opt.step()

        gen_opt.zero_grad()
        # Calculate generator loss
        g_loss = ____
        # Compute generator gradients
        ____
        gen_opt.step()

        print(f"Generator loss: {g_loss}")
        print(f"Discriminator loss: {d_loss}")
        break
Redigera och kör kod