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
Övningsinstruktioner
- Beräkna diskriminatorns förlust med
disc_loss()genom att skicka in generatorn, diskriminatorn, urvalet av riktiga bilder, aktuell batchstorlek och brusdimensionen16, i den ordningen, och tilldela resultatet tilld_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 brusdimensionen16, i den ordningen, och tilldela resultatet tillg_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