Pętla treningowa
Nadszedł czas, by cała praca włożona w definiowanie architektur modeli i funkcji straty przyniosła efekty – czas na trening! Twoim zadaniem jest zaimplementowanie i uruchomienie pętli treningowej GAN. Uwaga: po pierwszej partii danych umieszczono instrukcję break, aby uniknąć długiego czasu wykonania.
Oba optymizatory, disc_opt i gen_opt, zostały zainicjalizowane jako optymizatory Adam(). Funkcje do obliczania strat zdefiniowane wcześniej – gen_loss() i disc_loss() – są już dostępne. Przygotowano również dataloader.
Pamiętaj, że:
- Argumenty
disc_loss()to:gen,disc,real,cur_batch_size,z_dim. - Argumenty
gen_loss()to:gen,disc,cur_batch_size,z_dim.
To ćwiczenie jest częścią kursu
Głębokie uczenie dla obrazów z PyTorch
Instrukcje do ćwiczenia
- Oblicz stratę dyskryminatora za pomocą
disc_loss(), przekazując jej generator, dyskryminator, próbkę prawdziwych obrazów, bieżący rozmiar partii oraz rozmiar szumu16– w tej kolejności – i przypisz wynik dod_loss. - Oblicz gradienty, używając
d_loss. - Oblicz stratę generatora za pomocą
gen_loss(), przekazując jej generator, dyskryminator, bieżący rozmiar partii oraz rozmiar szumu16– w tej kolejności – i przypisz wynik dog_loss. - Oblicz gradienty, używając
g_loss.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
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