Entraîner un modèle GAN
Votre équipe chez PyBooks a bien avancé dans la création d'un générateur de texte avec un réseau antagoniste génératif (GAN). Vous avez défini avec succès les réseaux du générateur et du discriminateur. Il est maintenant temps de les entraîner. L'étape finale consiste à générer des fausses données et à les comparer aux vraies données pour voir dans quelle mesure votre GAN a appris. Nous avons utilisé des tenseurs comme entrée et la sortie tentera de ressembler aux tenseurs d'entrée. L'équipe de PyBooks pourra ensuite utiliser ces données synthétiques pour l'analyse de texte, car les caractéristiques auront les mêmes relations que les données textuelles.
Le générateur et le discriminateur ont été initialisés et enregistrés dans generator et discriminator, respectivement.
Les variables suivantes ont été initialisées dans l'exercice :
seq_length = 5: longueur de chaque séquence de données synthétiquesnum_sequences = 100: nombre total de séquences généréesnum_epochs = 50: nombre de passages complets dans l'ensemble de donnéesprint_every = 10: fréquence d'affichage des résultats, toutes les 10 époques
Cette activité fait partie du cours
Apprentissage profond pour le texte avec PyTorch
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
# 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()