Strata generatora
Zanim zaczniesz trenować sieć GAN, musisz zdefiniować funkcje straty dla generatora i dyskryminatora. Zacznij od tego pierwszego.
Przypomnij sobie, że zadaniem generatora jest tworzenie fałszywych obrazów, które zmylą dyskryminatora i skłonią go do sklasyfikowania ich jako prawdziwe. Generator ponosi więc stratę, gdy dyskryminator rozpoznaje wygenerowane obrazy jako fałszywe (etykieta 0).
Zdefiniuj funkcję gen_loss(), która oblicza stratę generatora. Przyjmuje cztery argumenty:
gen– model generatoradisc– model dyskryminatoranum_images– liczba obrazów w partiiz_dim– rozmiar wejściowego losowego szumu
To ćwiczenie jest częścią kursu
Głębokie uczenie dla obrazów z PyTorch
Instrukcje do ćwiczenia
- Wygeneruj losowy szum o kształcie
num_imagesnaz_dimi przypisz go donoise. - Użyj generatora, aby wygenerować fałszywy obraz na podstawie
noisei przypisz go dofake. - Uzyskaj predykcję dyskryminatora dla wygenerowanego fałszywego obrazu.
- Oblicz stratę generatora, wywołując
criterionna predykcjach dyskryminatora i tensorze jedynek o tym samym kształcie.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
def gen_loss(gen, disc, criterion, num_images, z_dim):
# Define random noise
noise = ____(num_images, z_dim)
# Generate fake image
fake = ____
# Get discriminator's prediction on the fake image
disc_pred = ____
# Compute generator loss
criterion = nn.BCEWithLogitsLoss()
gen_loss = ____(____, ____)
return gen_loss