Zacznij terazZacznij za darmo

Trenowanie modelu GAN

Twój zespół w PyBooks poczynił duże postępy w budowie generatora tekstu opartego na generatywnej sieci przeciwstawnej (GAN). Udało się zdefiniować sieci generatora i dyskryminatora. Teraz czas je wytrenować. Ostatni krok to wygenerowanie fałszywych danych i porównanie ich z danymi rzeczywistymi, aby sprawdzić, jak dobrze model GAN się nauczył. Jako dane wejściowe użyto tensorów, a dane wyjściowe będą starały się je naśladować. Zespół PyBooks będzie mógł następnie wykorzystać te syntetyczne dane do analizy tekstu – cechy będą miały takie same relacje jak w danych tekstowych.

Generator i dyskryminator zostały zainicjalizowane i zapisane odpowiednio do zmiennych generator i discriminator.

Następujące zmienne zostały zainicjalizowane w ćwiczeniu:

  • seq_length = 5: długość każdej sekwencji syntetycznych danych
  • num_sequences = 100: łączna liczba wygenerowanych sekwencji
  • num_epochs = 50: liczba pełnych przebiegów przez zbiór danych
  • print_every = 10: częstotliwość wyświetlania wyników – co 10 epok

To ćwiczenie jest częścią kursu

Uczenie głębokie dla tekstu z PyTorch

Zobacz kurs

Interaktywne ćwiczenie praktyczne

Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.

# Define the loss function and optimizer
criterion = nn.____()
optimizer_gen = ____(generator.parameters(), lr=0.001)
optimizer_disc = ____(discriminator.parameters(), lr=0.001)

for epoch in range(num_epochs):
    for real_data in data:
      	# Unsqueezing real_data and prevent gradient recalculations
        real_data = real_data.____(0)
        noise = torch.rand((1, seq_length))
        fake_data = generator(noise)
        disc_real = discriminator(real_data)
        disc_fake = discriminator(fake_data.____())
        loss_disc = criterion(disc_real, torch.ones_like(disc_real)) + criterion(disc_fake, torch.zeros_like(disc_fake))
        optimizer_disc.zero_grad()
        loss_disc.backward()
        optimizer_disc.step()
Edytuj i uruchom kod