Entraîner des GAN

Deep Learning pour les images avec PyTorch

Michal Oleszak

Machine Learning Engineer

Objectif du générateur

 

 

Schéma du flux de travail d'un GAN.

  • Objectif : générer des faux qui trompent le discriminateur
  • Idée : s'appuyer sur le discriminateur pour évaluer le générateur
  • Sortie du générateur classée par le discriminateur :
    • Réelle (étiquette 1) : bien, faible perte
    • Fausse (étiquette 0) : mal, grande perte
Deep Learning pour les images avec PyTorch

Perte du générateur

def gen_loss(gen, disc, num_images, z_dim):

noise = torch.randn(num_images, z_dim)
fake = gen(noise)
disc_pred = disc(fake)
criterion = nn.BCEWithLogitsLoss()
gen_loss = criterion( disc_pred, torch.ones_like(disc_pred) ) return gen_loss
  • Définir un bruit aléatoire
  • Générer une image fausse
  • Obtenir la prédiction du discriminateur sur l'image fausse
  • Utiliser l'entropie croisée binaire (BCE) comme critère
  • Perte du générateur : BCE entre les prédictions du discriminateur et un tenseur de uns
Deep Learning pour les images avec PyTorch

Objectif du discriminateur

 

 

Schéma du flux de travail d'un GAN.

  • Objectif : bien classer les faux et les images réelles
  • Les sorties du générateur doivent être classées comme fausses (étiquette 0)
  • Les images réelles doivent être classées comme réelles (étiquette 1)
Deep Learning pour les images avec PyTorch

Perte du discriminateur

def disc_loss(gen, disc, real, num_images, z_dim):

criterion = nn.BCEWithLogitsLoss()
noise = torch.randn(num_images, z_dim)
fake = gen(noise)
disc_pred_fake = disc(fake)
fake_loss = criterion( disc_pred_fake, torch.zeros_like(disc_pred_fake) )
disc_pred_real = disc(real)
real_loss = criterion( disc_pred_real, torch.ones_like(disc_pred_real) )
disc_loss = (real_loss + fake_loss) / 2 return disc_loss
  • Définir le critère d'entropie croisée binaire
  • Générer le bruit d'entrée du générateur
  • Générer des faux
  • Obtenir les prédictions du discriminateur pour les images fausses
  • Calculer la composante de perte pour les faux
  • Obtenir les prédictions du discriminateur pour les images réelles
  • Calculer la composante de perte pour les réelles
  • La perte finale est la moyenne des composantes « réelles » et « fausses »
Deep Learning pour les images avec PyTorch

Boucle d'entraînement du GAN

for epoch in range(num_epochs):
    for real in dataloader:
        cur_batch_size = len(real)


disc_opt.zero_grad()
disc_loss = disc_loss( gen, disc, real, cur_batch_size, z_dim=16)
disc_loss.backward() disc_opt.step()
gen_opt.zero_grad()
gen_loss = gen_loss( gen, disc, cur_batch_size, z_dim=16)
gen_loss.backward() gen_opt.step()
  • Boucler sur les époques et les lots d'images réelles, puis calculer la taille courante du lot
  • Réinitialiser les gradients de l'optimiseur du discriminateur
  • Calculer la perte du discriminateur
  • Calculer les gradients du discriminateur et effectuer l'étape d'optimisation
  • Réinitialiser les gradients de l'optimiseur du générateur
  • Calculer la perte du générateur
  • Calculer les gradients du générateur et effectuer l'étape d'optimisation
Deep Learning pour les images avec PyTorch

Passons à la pratique !

Deep Learning pour les images avec PyTorch

Preparing Video For Download...