CelebA GAN
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:
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
# 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)
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...
É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.
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
Epoch 10
Évolution (epoch 1 → 10)
É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:
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
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
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.
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
É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:
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)))
RÉSULTATS: