Strata dyskryminatora
Czas zdefiniować stratę dyskryminatora. Przypomnij sobie, że zadaniem dyskryminatora jest klasyfikowanie obrazów jako prawdziwe lub fałszywe. Dyskryminator ponosi stratę, gdy sklasyfikuje wyniki generatora jako prawdziwe (etykieta 1) lub gdy uzna prawdziwe obrazy za fałszywe (etykieta 0).
Zdefiniuj funkcję disc_loss(), która oblicza stratę dyskryminatora. Przyjmuje ona pięć argumentów:
gen– model generatoradisc– model dyskryminatorareal– próbka prawdziwych obrazów ze zbioru treningowegonum_images– liczba obrazów w partiiz_dim– rozmiar wejściowego szumu losowego
To ćwiczenie jest częścią kursu
Głębokie uczenie dla obrazów z PyTorch
Instrukcje do ćwiczenia
- Użyj dyskryminatora do sklasyfikowania obrazów
fakei przypisz przewidywania do zmiennejdisc_pred_fake. - Oblicz składnik straty dla fałszywych obrazów, wywołując
criterionna przewidywaniach dyskryminatora dla fałszywych obrazów oraz tensora samych zer o tym samym kształcie. - Użyj dyskryminatora do sklasyfikowania obrazów
reali przypisz przewidywania do zmiennejdisc_pred_real. - Oblicz składnik straty dla prawdziwych obrazów, wywołując
criterionna przewidywaniach dyskryminatora dla prawdziwych obrazów oraz tensora samych jedynek o tym samym kształcie.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
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