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 danychnum_sequences = 100: łączna liczba wygenerowanych sekwencjinum_epochs = 50: liczba pełnych przebiegów przez zbiór danychprint_every = 10: częstotliwość wyświetlania wyników – co 10 epok
To ćwiczenie jest częścią kursu
Uczenie głębokie dla tekstu z PyTorch
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()