mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-06-25 07:59:55 +02:00
Introduce new AdversarialLoss class
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import configparser
|
||||
from typing import Tuple
|
||||
from typing import List, Tuple
|
||||
from warnings import deprecated
|
||||
|
||||
import torch
|
||||
from pytorch_msssim import ssim
|
||||
@@ -78,6 +79,7 @@ class FaceSwapperLoss:
|
||||
discriminator_loss_set['loss_discriminator'] = (loss_true + loss_fake) * 0.5
|
||||
return discriminator_loss_set
|
||||
|
||||
@deprecated
|
||||
def calc_adversarial_loss(self, discriminator_outputs : DiscriminatorOutputs) -> LossTensor:
|
||||
loss_adversarials = []
|
||||
|
||||
@@ -96,6 +98,7 @@ class FaceSwapperLoss:
|
||||
loss_attribute = torch.stack(loss_attributes).mean() * 0.5
|
||||
return loss_attribute
|
||||
|
||||
@deprecated
|
||||
def calc_reconstruction_loss(self, swap_tensor : VisionTensor, target_tensor : VisionTensor, is_same_person : Tensor) -> LossTensor:
|
||||
loss_reconstruction = torch.pow(swap_tensor - target_tensor, 2).reshape(self.batch_size, -1)
|
||||
loss_reconstruction = torch.mean(loss_reconstruction, dim = 1) * 0.5
|
||||
@@ -104,6 +107,7 @@ class FaceSwapperLoss:
|
||||
loss_reconstruction = (loss_reconstruction + loss_ssim) * 0.5
|
||||
return loss_reconstruction
|
||||
|
||||
@deprecated
|
||||
def calc_identity_loss(self, source_tensor : VisionTensor, swap_tensor : VisionTensor) -> LossTensor:
|
||||
swap_embedding = calc_embedding(self.embedder, swap_tensor, (30, 0, 10, 10))
|
||||
source_embedding = calc_embedding(self.embedder, source_tensor, (30, 0, 10, 10))
|
||||
@@ -141,25 +145,44 @@ class FaceSwapperLoss:
|
||||
return translation, scale, rotation
|
||||
|
||||
|
||||
class AdversarialLoss(torch.nn.Module):
|
||||
def __init__(self) -> None:
|
||||
super(AdversarialLoss, self).__init__()
|
||||
|
||||
def calc(self, discriminator_output_tensors : List[Tensor]) -> Tuple[Tensor, Tensor]:
|
||||
adversarial_weight = CONFIG.getfloat('training.losses', 'adversarial_weight')
|
||||
temp_tensors = []
|
||||
|
||||
for discriminator_output_tensor in discriminator_output_tensors:
|
||||
temp_tensor = torch.relu(1 - discriminator_output_tensor[0]).mean()
|
||||
temp_tensors.append(temp_tensor)
|
||||
|
||||
loss = torch.stack(temp_tensors).mean()
|
||||
weighted_loss = loss * adversarial_weight
|
||||
return loss, weighted_loss
|
||||
|
||||
|
||||
class ReconstructionLoss(torch.nn.Module):
|
||||
def __init__(self) -> None:
|
||||
super(ReconstructionLoss, self).__init__()
|
||||
|
||||
def calc(self, source_tensor : Tensor, target_tensor : Tensor, output_tensor : Tensor) -> Tensor:
|
||||
def calc(self, source_tensor : Tensor, target_tensor : Tensor, output_tensor : Tensor) -> Tuple[Tensor, Tensor]:
|
||||
batch_size = CONFIG.getint('training.loader', 'batch_size')
|
||||
loss_tensor = torch.pow(output_tensor - target_tensor, 2).reshape(batch_size, -1)
|
||||
loss_tensor = torch.mean(loss_tensor, dim = 1) * 0.5
|
||||
reconstruction_weight = CONFIG.getfloat('training.losses', 'reconstruction_weight')
|
||||
loss = torch.pow(output_tensor - target_tensor, 2).reshape(batch_size, -1)
|
||||
loss = torch.mean(loss, dim = 1) * 0.5
|
||||
|
||||
if torch.equal(source_tensor, target_tensor):
|
||||
loss_tensor = torch.sum(loss_tensor * torch.tensor(0)) / (torch.tensor(0).sum() + 1e-4)
|
||||
loss = torch.sum(loss * torch.tensor(0)) / (torch.tensor(0).sum() + 1e-4)
|
||||
else:
|
||||
loss_tensor = torch.sum(loss_tensor * torch.tensor(1)) / (torch.tensor(1).sum() + 1e-4)
|
||||
loss = torch.sum(loss * torch.tensor(1)) / (torch.tensor(1).sum() + 1e-4)
|
||||
|
||||
data_range = float(torch.max(output_tensor) - torch.min(output_tensor))
|
||||
similarity = 1 - ssim(output_tensor, target_tensor, data_range = data_range).mean()
|
||||
|
||||
loss_tensor = (loss_tensor + similarity) * 0.5
|
||||
return loss_tensor
|
||||
loss = (loss + similarity) * 0.5
|
||||
weighted_loss = loss * reconstruction_weight
|
||||
return loss, weighted_loss
|
||||
|
||||
|
||||
class IdentityLoss(torch.nn.Module):
|
||||
@@ -169,8 +192,10 @@ class IdentityLoss(torch.nn.Module):
|
||||
self.embedder = torch.jit.load(embedder_path, map_location = 'cpu') # type:ignore[no-untyped-call]
|
||||
self.embedder.eval()
|
||||
|
||||
def calc(self, source_tensor : Tensor, output_tensor : Tensor) -> Tensor:
|
||||
def calc(self, source_tensor : Tensor, output_tensor : Tensor) -> Tuple[Tensor, Tensor]:
|
||||
identity_weight = CONFIG.getfloat('training.losses', 'identity_weight')
|
||||
output_embedding = calc_embedding(self.embedder, output_tensor, (30, 0, 10, 10))
|
||||
source_embedding = calc_embedding(self.embedder, source_tensor, (30, 0, 10, 10))
|
||||
loss_tensor = (1 - torch.cosine_similarity(source_embedding, output_embedding)).mean()
|
||||
return loss_tensor
|
||||
loss = (1 - torch.cosine_similarity(source_embedding, output_embedding)).mean()
|
||||
weighted_loss = loss * identity_weight
|
||||
return loss, weighted_loss
|
||||
|
||||
@@ -16,7 +16,7 @@ from .dataset import DynamicDataset
|
||||
from .helper import calc_embedding
|
||||
from .models.discriminator import Discriminator
|
||||
from .models.generator import Generator
|
||||
from .models.loss import FaceSwapperLoss, IdentityLoss, ReconstructionLoss
|
||||
from .models.loss import AdversarialLoss, FaceSwapperLoss, IdentityLoss, ReconstructionLoss
|
||||
from .types import Batch, Embedding, VisionTensor
|
||||
|
||||
CONFIG = configparser.ConfigParser()
|
||||
@@ -31,6 +31,7 @@ class FaceSwapperTrainer(lightning.LightningModule, FaceSwapperLoss):
|
||||
|
||||
self.generator = Generator()
|
||||
self.discriminator = Discriminator()
|
||||
self.adversarial_loss = AdversarialLoss()
|
||||
self.reconstruction_loss = ReconstructionLoss()
|
||||
self.identity_loss = IdentityLoss()
|
||||
self.automatic_optimization = automatic_optimization
|
||||
@@ -54,17 +55,17 @@ class FaceSwapperTrainer(lightning.LightningModule, FaceSwapperLoss):
|
||||
target_attributes = self.generator.get_attributes(target_tensor)
|
||||
generator_output_tensor = self.generator(source_embedding, target_tensor)
|
||||
generator_output_attributes = self.generator.get_attributes(generator_output_tensor)
|
||||
discriminator_output_tensor = self.discriminator(generator_output_tensor)
|
||||
discriminator_output_tensors = self.discriminator(generator_output_tensor)
|
||||
|
||||
generator_loss_set = self.calc_generator_loss(generator_output_tensor, target_attributes, generator_output_attributes, discriminator_output_tensor, batch)
|
||||
generator_loss_set = self.calc_generator_loss(generator_output_tensor, target_attributes, generator_output_attributes, discriminator_output_tensors, batch)
|
||||
generator_optimizer.zero_grad()
|
||||
self.manual_backward(generator_loss_set.get('loss_generator'))
|
||||
generator_optimizer.step()
|
||||
|
||||
discriminator_source_tensor = self.discriminator(source_tensor)
|
||||
discriminator_output_tensor = self.discriminator(generator_output_tensor.detach())
|
||||
discriminator_output_tensors = self.discriminator(generator_output_tensor.detach())
|
||||
|
||||
discriminator_loss_set = self.calc_discriminator_loss(discriminator_source_tensor, discriminator_output_tensor)
|
||||
discriminator_loss_set = self.calc_discriminator_loss(discriminator_source_tensor, discriminator_output_tensors)
|
||||
discriminator_optimizer.zero_grad()
|
||||
self.manual_backward(discriminator_loss_set.get('loss_discriminator'))
|
||||
discriminator_optimizer.step()
|
||||
@@ -74,30 +75,24 @@ class FaceSwapperTrainer(lightning.LightningModule, FaceSwapperLoss):
|
||||
|
||||
self.log('loss_generator', generator_loss_set.get('loss_generator'), prog_bar = True)
|
||||
self.log('loss_discriminator', discriminator_loss_set.get('loss_discriminator'), prog_bar = True)
|
||||
self.log('loss_adversarial', generator_loss_set.get('loss_adversarial'))
|
||||
self.log('loss_adversarial', generator_loss_set.get('loss_adversarial'), prog_bar = True)
|
||||
self.log('loss_attribute', generator_loss_set.get('loss_attribute'))
|
||||
self.log('loss_identity', generator_loss_set.get('loss_identity'), prog_bar = True)
|
||||
self.log('loss_reconstruction', generator_loss_set.get('loss_reconstruction'), prog_bar = True)
|
||||
self.log('loss_identity', generator_loss_set.get('loss_identity'))
|
||||
self.log('loss_reconstruction', generator_loss_set.get('loss_reconstruction'))
|
||||
|
||||
reconstruction_loss = self.reconstruction_loss.calc(source_tensor, target_tensor, generator_output_tensor)
|
||||
identity_loss = self.identity_loss.calc(generator_output_tensor, source_tensor)
|
||||
generator_loss = self.calc_generator_loss_new(reconstruction_loss, identity_loss)
|
||||
###############################################
|
||||
|
||||
adversarial_loss, weighted_adversarial_loss = self.adversarial_loss.calc(discriminator_output_tensors)
|
||||
reconstruction_loss, weighted_reconstruction_loss = self.reconstruction_loss.calc(source_tensor, target_tensor, generator_output_tensor)
|
||||
identity_loss, weighted_identity_loss = self.identity_loss.calc(generator_output_tensor, source_tensor)
|
||||
generator_loss = weighted_adversarial_loss+ weighted_reconstruction_loss + weighted_identity_loss
|
||||
|
||||
self.log('generator_loss_new', generator_loss, prog_bar = True)
|
||||
self.log('loss_reconstruction_new', reconstruction_loss, prog_bar = True)
|
||||
self.log('loss_identity_new', identity_loss, prog_bar = True)
|
||||
self.log('adversarial_loss_new', adversarial_loss, prog_bar = True)
|
||||
self.log('loss_reconstruction_new', reconstruction_loss)
|
||||
self.log('loss_identity_new', identity_loss)
|
||||
return generator_loss_set.get('loss_generator')
|
||||
|
||||
@staticmethod
|
||||
def calc_generator_loss_new(reconstruction_loss : Tensor, identity_loss : Tensor) -> Tensor:
|
||||
reconstruction_weight = CONFIG.getfloat('training.losses', 'reconstruction_weight')
|
||||
identity_weight = CONFIG.getfloat('training.losses', 'identity_weight')
|
||||
|
||||
generator_loss = reconstruction_loss * reconstruction_weight
|
||||
generator_loss += identity_loss * identity_weight
|
||||
|
||||
return generator_loss
|
||||
|
||||
def validation_step(self, batch : Batch, batch_index : int) -> Tensor:
|
||||
source_tensor, target_tensor = batch
|
||||
source_embedding = calc_embedding(self.embedder, source_tensor, (0, 0, 0, 0))
|
||||
|
||||
Reference in New Issue
Block a user