Zacznij terazZacznij za darmo

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 generatora
  • disc – model dyskryminatora
  • real – próbka prawdziwych obrazów ze zbioru treningowego
  • num_images – liczba obrazów w partii
  • z_dim – rozmiar wejściowego szumu losowego

To ćwiczenie jest częścią kursu

Głębokie uczenie dla obrazów z PyTorch

Zobacz kurs

Instrukcje do ćwiczenia

  • Użyj dyskryminatora do sklasyfikowania obrazów fake i przypisz przewidywania do zmiennej disc_pred_fake.
  • Oblicz składnik straty dla fałszywych obrazów, wywołując criterion na przewidywaniach dyskryminatora dla fałszywych obrazów oraz tensora samych zer o tym samym kształcie.
  • Użyj dyskryminatora do sklasyfikowania obrazów real i przypisz przewidywania do zmiennej disc_pred_real.
  • Oblicz składnik straty dla prawdziwych obrazów, wywołując criterion na 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
Edytuj i uruchom kod