CelebA GAN

gan architecture

Sur cette page je vais expliquer comment j'ai entrainé mon propre modèle d'IA dans le but de faire des deepfake à partir du dataset CelebA

Objectif

Ce n'est pas la première fois que je realise ce projet, j'ai déjà tenté l'expérience en 2022 mais je m'étais confronté à plusieurs problèmes:

  • Je n'avais aucune idée de comment mesurer les progrès de mon modèle, à cause des fonctionnement des GAN, il n'y a pas vraiment de loss qui diminue
  • En conséquent, je n'avais aucune idée de combien de compute je devais allouer, parle-t-on d'heures ? De jours ? De semaine ?
  • J'étais étudiant et j'avais peur d'allouer un budget, aujourd'hui je peux me permettre de louer une VM sur google cloud si besoin, je peux donc tester des modèles plus gros et plus longtemps
  • Mon objectif aujourd'hui est de reproduire l'expérience avec un plus gros budget, de meilleurs connaissances et un état de l'art plus poussé.

    ÉTAPE 1: Structure de base

    model.py
    
    # Generator model, takes a random seed vector and make an image
    # It is trained to fool the Discriminator.
    class Discriminator(torch.nn.Module):
        def __init__(self):
            super().__init__()
            self.net = torch.nn.Sequential(
                torch.nn.Conv2d(3, 64, 3, stride=2, padding=1),    torch.nn.LeakyReLU(0.2),
                torch.nn.Conv2d(64, 128, 3, stride=2, padding=1),  torch.nn.LeakyReLU(0.2),
                torch.nn.Conv2d(128, 256, 3, stride=2, padding=1), torch.nn.LeakyReLU(0.2),
                torch.nn.Flatten(), torch.nn.Dropout(0.4), torch.nn.Linear(flat_size, 1),
            )
    
        def forward(self, x):
            return self.net(x)
    
    # The Discriminator is trained to figure out if an image was made by the Generator or not.
    class Generator(torch.nn.Module):
        def __init__(self, latent_size):
            super().__init__()
            self.net = torch.nn.Sequential(
                torch.nn.Linear(latent_size, 64 * 28 * 32), torch.nn.LeakyReLU(0.2),
                Reshape(64, 28, 32),
                torch.nn.ConvTranspose2d(64, 128, 4, stride=2, padding=1), torch.nn.LeakyReLU(0.2),
                torch.nn.ConvTranspose2d(128, 3, 4, stride=2, padding=1), torch.nn.Tanh(),
            )
    
        def forward(self, x):
            return self.net(x)
        
    core.py
    
    def train_step(optimizer, loss):
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
    
    
    def train_epoch(g_model, d_model, g_optimizer, d_optimizer, dataset, config):
        # We use Binary Cross Entropy loss, standard for True/False prediction
        criterion = torch.nn.BCELoss()
        g_model.train(); d_model.train()
    
        for batch, _ in dataset:
            batch = batch.cuda()
            real_labels = torch.ones(len(batch), 1).cuda()
            fake_labels = torch.zeros(len(batch), 1).cuda()
    
            train_step(d_optimizer, criterion(d_model(batch), real_labels))
            train_step(d_optimizer, criterion(d_model(generate_fake(g_model, config)), fake_labels))
    
            # Last time I tested, the Generator was training slower than the Discriminator, this is expected:
            # It's probably way easier to discriminate an image than to generate a convincing one.
            # However, if the Discriminator gets too good and guess right every time, the Generator cannot
            # learn what's working. To balance this, I decided to train the generator twice, but we will find
            # more elegant solutions later.
            for _ in range(2):
                train_step(g_optimizer, criterion(d_model(generate_fake(g_model, config)), real_labels))
    
    
    def train(config):
        g_model, d_model, g_optimizer, d_optimizer = make_model(config)
        dataset = make_dataset(config)
    
        for epoch in range(config["number_epochs"]):
            update_lr(g_optimizer, d_optimizer, config)
            train_epoch(g_model, d_model, g_optimizer, d_optimizer, dataset, config)
            log_samples(g_model, epoch)
    
        save_model(g_model, d_model, config)
        

    RÉSULTATS:

    Résultats après une (1) epoch, c'est-à-dire un entrainement sur chacune des images une fois. Le but de cet entrainement est surtout de valider le code, tester l'intégration avec Wandb, s'assurer que la loss ne dégénère pas...
    GAN sample 1 GAN sample 2 GAN sample 3 GAN sample 4 GAN sample 5

    ÉTAPE 2: L'architecture des modèles

    Le papier original de DCGAN (DCGAN étant les GAN avec des couches de convolution) propose 4 couches de convolution, comme en IA personne ne sait ce qu'il fait et qu'on se contente de faire ce qui marche parce que ça marche je vais prendre une architecture similaire.
    model.py
    
    class Discriminator(torch.nn.Module): # 1.58M parameters
        def __init__(self):
            super().__init__()
            self.features = torch.nn.Sequential(
                torch.nn.utils.spectral_norm(
                    torch.nn.Conv2d(3, 64, kernel_size=3, stride=2, padding=1)
                ), # 3 * 64 * 3 * 3 + 64 = 1,792
                torch.nn.LeakyReLU(0.2),
                torch.nn.utils.spectral_norm(
                    torch.nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1)
                ), # 64 * 128 * 3 * 3 + 128 = 73,856
                torch.nn.LeakyReLU(0.2),
                torch.nn.utils.spectral_norm(
                    torch.nn.Conv2d(128, 256, kernel_size=3, stride=2, padding=1)
                ), # 128 * 256 * 3 * 3 + 256 = 295,168
                torch.nn.LeakyReLU(0.2),
                torch.nn.utils.spectral_norm(
                    torch.nn.Conv2d(256, 512, kernel_size=3, stride=2, padding=1)
                ), # 256 * 512 * 3 * 3 + 512 = 1,180,160
                torch.nn.LeakyReLU(0.2),
                torch.nn.Flatten(),
            )
            # 4 stride-2 convs on 112x128: 112->56->28->14->7, 128->64->32->16->8
            # flat_size = 512 * 7 * 8 = 28,672
            dummy = torch.zeros(1, *INPUT_SHAPE)
            flat_size = self.features(dummy).shape[1]
            self.classifier = torch.nn.Sequential(
                torch.nn.Dropout(0.4),
                torch.nn.Linear(flat_size, 1), # 28,672 * 1 + 1 = 28,673
            )
    
        def forward(self, x):
            return self.classifier(self.features(x))
    
    
    class Generator(torch.nn.Module): # 5.69M parameters
        def __init__(self, latent_size: int):
            super().__init__()
            n_nodes = 512 * (112 // (2**4)) * (128 // (2**4))  # 512 * 7 * 8 = 28,672
            self.features = torch.nn.Sequential(
                torch.nn.Linear(latent_size, n_nodes),              # latent_size * 28,672 + 28,672
                                                                    # e.g. latent=100: 2,895,872
                torch.nn.LeakyReLU(0.2),
                Reshape(512, 7, 8),
                torch.nn.ConvTranspose2d(512, 256, kernel_size=4, stride=2, padding=1),  # 512 * 256 * 4 * 4 + 256 = 2,097,408
                torch.nn.BatchNorm2d(256),                          # 256 * 2 = 512
                torch.nn.LeakyReLU(0.2),
                torch.nn.ConvTranspose2d(256, 128, kernel_size=4, stride=2, padding=1),  # 256 * 128 * 4 * 4 + 128 = 524,416
                torch.nn.BatchNorm2d(128),                          # 128 * 2 = 256
                torch.nn.LeakyReLU(0.2),
                torch.nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1),   # 128 * 64 * 4 * 4 + 64 = 131,136
                torch.nn.BatchNorm2d(64),                           # 64 * 2 = 128
                torch.nn.LeakyReLU(0.2),
                torch.nn.ConvTranspose2d(64, 32, kernel_size=4, stride=2, padding=1),    # 64 * 32 * 4 * 4 + 32 = 32,800
                torch.nn.BatchNorm2d(32),                           # 32 * 2 = 64
                torch.nn.LeakyReLU(0.2),
            )
            self.generator = torch.nn.Sequential(
                torch.nn.ConvTranspose2d(32, 3, kernel_size=7, padding=3),               # 32 * 3 * 7 * 7 + 3 = 4,707
                torch.nn.Tanh()
            )
    
        def forward(self, x):
            return self.generator(self.features(x))
    

    RÉSULTATS:

    Cette fois, je m'autorise à aller jusqu'à 10 epochs. Pour rappel, le Generator génère des images à partir d'un vecteur seed aléatoire. Cependant, si on fixe un vecteur seed, on peut voir l'évolution du modèle sur le même "prompt".

    Epoch 1

    Sample 1 epoch 1 Sample 2 epoch 1 Sample 3 epoch 1 Sample 4 epoch 1 Sample 5 epoch 1
    ↓ ↓ ↓ ↓ ↓

    Epoch 10

    Sample 1 epoch 10 Sample 2 epoch 10 Sample 3 epoch 10 Sample 4 epoch 10 Sample 5 epoch 10

    Évolution (epoch 1 → 10)

    Evolution 1 Evolution 2 Evolution 3 Evolution 4 Evolution 5

    ÉTAPE 3: Équilibrage D/G

    Durant l'étape 2, il y avait toujours un écart de loss, entre le Discriminator et le Generator, pour parer à cela, nous avons plusieurs options à notre disposition:
  • Label Smoothing: Cela consiste à dire qu'une vraie image n'est pas vraie à 100%, mais seulement à quelque chose autours de 95%, le but étant de gaslight le discriminateur pour qu'il ne prène pas trop la confiance.
  • Adaptative Learning Rate: La version propre de ma première solution (double entrainement pour le Generator), on peut simplement comparer les Loss, et ponderer le learning rate en fonction du delta
  • Voici un exemple d'adaptative learining rate:
    model.py
    
    def compute_lr_scale(d_loss_real, d_loss_fake, g_loss, base_lr_d, base_lr_g, strength=0.5):
        LN2 = math.log(2)  # ~0.693, equilibrium loss value
    
        # D is winning when EITHER term is too low
        d_best = min(d_loss_real, d_loss_fake)
    
        d_advantage = LN2 - d_best  # large when D is crushing fakes
        # G is winning when its loss is low
        g_advantage = LN2 - g_loss
    
        # scale lr down for whoever is ahead, clamp to [0.1, 1.0]
        d_scale = (1.0 - strength * max(0, d_advantage) / LN2)
        g_scale = (1.0 - strength * max(0, g_advantage) / LN2)
    
        d_scale = max(0.1, min(1.0, d_scale))
        g_scale = max(0.1, min(1.0, g_scale))
    
        return base_lr_d * d_scale, base_lr_g * g_scale
    
    Après une run de 50 epochs, on a un total model collapse avec la loss du Generator qui monte encore très haut. Le but du jeu ça va être de tester des formules au pif jusqu'à ce qu'on trouve un truc qui marche
    core.py
    
    def compute_lr_scale(d_loss_real, d_loss_fake, g_loss, base_lr_d, base_lr_g, strength=0.5):
        LN2 = math.log(2)  # ~0.693, equilibrium loss value
    
        # D is winning when EITHER term is too low
        d_best = min(d_loss_real, d_loss_fake)
    
        d_advantage = LN2 - d_loss_fake  # Actually it's fine if d_loss_real is high, it is independent of our Generator
        # G is winning when its loss is low
        g_advantage = LN2 - g_loss
    
        # scale lr down for whoever is ahead, clamp to [0.1, 1.0]
        d_scale = (1.0 - strength * max(0, d_advantage) / LN2)
        g_scale = (1.0 - strength * max(0, g_advantage) / LN2)
    
        d_scale = max(0.1, min(1.0, d_scale))
        g_scale = max(0.1, min(1.0, g_scale))
    
        return base_lr_d * d_scale, base_lr_g * g_scale
    
    Une autre idée, c'est d'utiliser la pénalité R1, qui va pénaliser le Discriminator s'il est sensible à de petites variations:

    En principe, une image qu'elle soit vraie ou fausse, ne change pas de label pour un pixel, ou une petite variation. La pénalité R1 augmente la loss du Discriminator s'il est sensible aux petites variations.

    core.py
    
    
    def train_epoch(g_model, d_model, g_optimizer, d_optimizer, g_scaler, d_scaler, dataset, model_config, training_config, num_epoch, global_step):
    rength=0.5):
      # R1 gradient penalty every r1_batch_mod steps
      if global_step % training_config["r1_batch_mod"] == 0:
          grad = torch.autograd.grad(
              d_real_out.sum(), batch, create_graph=True
          )[0]
          r1 = grad.pow(2).flatten(1).sum(1).mean()
          d_loss = d_loss_real + d_loss_fake + (10.0/2) * 16 * r1
      else:
          d_loss = d_loss_real + d_loss_fake
    
    
    Mettre en place ces outils d'équilibrage a été suffisant pour garder les losses autours de ln(2) run 5 losses
    Evolution 1 Evolution 2 Evolution 3 Evolution 4 Evolution 5

    ÉTAPE 4: Scaling

    Maintenant que tout marche, l'idée est de tenter sur des modèles plus gros, et sur un plus long entrainement:
    model.py
    
    class Discriminator(torch.nn.Module): # 2.38M parameters
        @staticmethod
        def down_block(dims_in, dims_out): # 9 * dims_out * (dims_in + dims_out) + 2 * dims_out
            return torch.nn.Sequential(
                torch.nn.utils.spectral_norm(
                    torch.nn.Conv2d(dims_in, dims_out, kernel_size=3, stride=2, padding=1) # dims_in * dims_out * 9 + dims_out
                ),
                torch.nn.LeakyReLU(0.2),
                torch.nn.utils.spectral_norm(
                    torch.nn.Conv2d(dims_out, dims_out, kernel_size=3, padding=1) # dims_out^2 * 9 + dims_out
                ),
                torch.nn.LeakyReLU(0.2),
            )
    
        def __init__(self):
            super().__init__()
            self.input = torch.nn.Sequential(
                torch.nn.utils.spectral_norm(torch.nn.Conv2d(3, 64, kernel_size=3, padding=1)),
                torch.nn.LeakyReLU(0.2),
            )
            self.features = torch.nn.Sequential(
                self.down_block(64, 64),     # 74k
                self.down_block(64, 128),    # 221k
                self.down_block(128, 256),   # 885k
                self.down_block(256, 256),   # 1.18M
            )
    
            flat_size = 256 * 7 * 8  # 14336
            self.classifier = torch.nn.Sequential(
                torch.nn.Flatten(),
                torch.nn.Dropout(0.4),
                torch.nn.Linear(flat_size, 1),
            )
    
        def forward(self, x):
            return self.classifier(self.features(self.input(x)))
    
    
    class Reshape(torch.nn.Module):
        def __init__(self, *shape):
            super().__init__()
            self.shape = shape
    
        def forward(self, x):
            return x.view(x.size(0), *self.shape)
    
    
    class Generator(torch.nn.Module): # 4.3M parameters
        @staticmethod
        def up_block(dims_in, dims_out): # 9 * c_out * (c_in + c_out) + 6 * c_out
            return torch.nn.Sequential(
                torch.nn.Upsample(scale_factor=2, mode="nearest"),
                torch.nn.Conv2d(dims_in, dims_out, kernel_size=3, padding=1), # 9 * dims_in * dims_out + dims_out
                torch.nn.BatchNorm2d(dims_out), # 2 * c_out
                torch.nn.LeakyReLU(0.2),
                torch.nn.Conv2d(dims_out, dims_out, kernel_size=3, padding=1), # 9 * dims_out^2 + dims_out
                torch.nn.BatchNorm2d(dims_out), # 2 * c_out
                torch.nn.LeakyReLU(0.2)
            )
    
        def __init__(self, latent_size: int): # 4.3M
            super().__init__()
            n_nodes = 256 * (112 // (2**4)) * (128 // (2**4))  # 256 * 7 * 8 = 28,672
            self.projection = torch.nn.Sequential(
                torch.nn.Linear(latent_size, n_nodes),  # 1.43M
                Reshape(256, 7, 8),
                torch.nn.BatchNorm2d(256),
                torch.nn.LeakyReLU(0.2),
            )
    
            self.features = torch.nn.Sequential(
                self.up_block(256, 256), # 1.18M
                self.up_block(256, 256), # 1.18M
                self.up_block(256, 256), # 1.8M
                self.up_block(256, 128), # 444k
            )
            self.generator = torch.nn.Sequential(
                torch.nn.Conv2d(128, 64, kernel_size=3, padding=1), # 72k
                torch.nn.LeakyReLU(0.2),
                torch.nn.Conv2d(64, 3, kernel_size=1),
                torch.nn.Tanh(),
            )
    
        def forward(self, x):
            return self.generator(self.features(self.projection(x)))
    J'ai entrainé ce modèle sur 100 epochs sur une VM GCP, j'ai dû tweak quelques hyper-paramètres pour stabiliser l'entrainement, mais je suis rapidement tombé sur une run stable

    RÉSULTATS:

    Sample 12 Sample 14 Sample 17 Sample 7 Sample 9
    ↓ ↓ ↓ ↓ ↓
    Evolution 12 Evolution 14 Evolution 17 Evolution 7 Evolution 9