From 8f1f002c648b8381c86a987eba9f926a61515173 Mon Sep 17 00:00:00 2001 From: harisreedhar Date: Sat, 8 Mar 2025 16:05:07 +0530 Subject: [PATCH] add masknet --- face_swapper/README.md | 1 + face_swapper/config.ini | 1 + face_swapper/src/models/generator.py | 9 ++- face_swapper/src/models/loss.py | 31 +++++++++- face_swapper/src/networks/masknet.py | 90 ++++++++++++++++++++++++++++ face_swapper/src/training.py | 19 +++--- face_swapper/src/types.py | 1 + 7 files changed, 142 insertions(+), 10 deletions(-) create mode 100644 face_swapper/src/networks/masknet.py diff --git a/face_swapper/README.md b/face_swapper/README.md index 7772c88..b7bf47f 100644 --- a/face_swapper/README.md +++ b/face_swapper/README.md @@ -74,6 +74,7 @@ identity_weight = 20.0 gaze_weight = 0.0 pose_weight = 0.0 expression_weight = 0.0 +mask_weight = 1.0 ``` ``` diff --git a/face_swapper/config.ini b/face_swapper/config.ini index 0f3c702..dc20073 100644 --- a/face_swapper/config.ini +++ b/face_swapper/config.ini @@ -36,6 +36,7 @@ identity_weight = gaze_weight = pose_weight = expression_weight = +mask_weight = [training.trainer] learning_rate = diff --git a/face_swapper/src/models/generator.py b/face_swapper/src/models/generator.py index d02be0b..84a9303 100644 --- a/face_swapper/src/models/generator.py +++ b/face_swapper/src/models/generator.py @@ -1,8 +1,10 @@ from configparser import ConfigParser +from typing import Tuple from torch import Tensor, nn from ..networks.aad import AAD +from ..networks.masknet import MaskNet from ..networks.unet import UNet from ..types import Attributes, Embedding @@ -12,13 +14,16 @@ class Generator(nn.Module): super().__init__() self.encoder = UNet(config_parser) self.generator = AAD(config_parser) + self.masker = MaskNet(67, 1, 16) self.encoder.apply(init_weight) self.generator.apply(init_weight) + self.masker.apply(init_weight) - def forward(self, source_embedding : Embedding, target_tensor : Tensor) -> Tensor: + def forward(self, source_embedding : Embedding, target_tensor : Tensor) -> Tuple[Tensor, Tensor]: target_attributes = self.get_attributes(target_tensor) output_tensor = self.generator(source_embedding, target_attributes) - return output_tensor + mask_tensor = self.masker(target_tensor, target_attributes[-1]) + return output_tensor, mask_tensor def get_attributes(self, input_tensor : Tensor) -> Attributes: return self.encoder(input_tensor) diff --git a/face_swapper/src/models/loss.py b/face_swapper/src/models/loss.py index 52f049f..3341d81 100644 --- a/face_swapper/src/models/loss.py +++ b/face_swapper/src/models/loss.py @@ -7,7 +7,7 @@ from torch import Tensor, nn from torchvision import transforms from ..helper import calc_embedding -from ..types import Attributes, EmbedderModule, Gaze, GazerModule, MotionExtractorModule +from ..types import Attributes, EmbedderModule, Gaze, GazerModule, MotionExtractorModule, ParserModule class DiscriminatorLoss(nn.Module): @@ -175,3 +175,32 @@ class GazeLoss(nn.Module): with torch.no_grad(): pitch, yaw = self.gazer(crop_tensor) return pitch, yaw + + +class MaskLoss(nn.Module): + def __init__(self, config_parser : ConfigParser, parser : ParserModule) -> None: + super().__init__() + self.config_mask_weight = config_parser.getfloat('training.losses', 'mask_weight') + self.config_output_size = config_parser.getint('training.model.generator', 'output_size') + self.parser = parser + self.mse_loss = nn.MSELoss() + + def forward(self, target_tensor : Tensor, mask_tensor : Tensor) -> Tuple[Tensor, Tensor]: + target_mask = self.calc_mask(target_tensor) + target_mask = target_mask.view(-1, self.config_output_size, self.config_output_size) + mask_tensor = mask_tensor.view(-1, self.config_output_size, self.config_output_size) + mask_loss = self.mse_loss(target_mask, mask_tensor) + weighted_mask_loss = mask_loss * self.config_mask_weight + return mask_loss, weighted_mask_loss + + def calc_mask(self, target_tensor : Tensor) -> Tensor: + target_tensor = torch.nn.functional.interpolate(target_tensor, (512, 512), mode = 'bilinear') + face_indices = torch.tensor([ 1, 2, 3, 4, 5, 10, 11, 12, 13 ]).to(target_tensor.device) + + with torch.no_grad(): + output_tensor = self.parser(target_tensor)[0] + output_tensor = output_tensor.argmax(1) + output_tensor = torch.isin(output_tensor, face_indices).to(target_tensor.dtype) + output_tensor = output_tensor.view(-1, 1, 512, 512) + output_tensor = torch.nn.functional.interpolate(output_tensor, (self.config_output_size, self.config_output_size), mode = 'bilinear') + return output_tensor diff --git a/face_swapper/src/networks/masknet.py b/face_swapper/src/networks/masknet.py new file mode 100644 index 0000000..a7d83d4 --- /dev/null +++ b/face_swapper/src/networks/masknet.py @@ -0,0 +1,90 @@ +import torch +from torch import Tensor, nn + + +class MaskNet(nn.Module): + def __init__(self, input_channels : int, output_channels : int, base_channels : int): + super().__init__() + self.down_samples = self.create_down_samples(input_channels, base_channels) + self.up_samples = self.create_up_samples(base_channels) + self.bottleneck = ResBlock(base_channels * 4) + self.conv = nn.Conv2d(base_channels, output_channels, kernel_size = 1) + self.sigmoid = nn.Sigmoid() + + def create_down_samples(self, input_channels : int, base_channels: int) -> nn.ModuleList: + down_samples = nn.ModuleList( + [ + DownSample(input_channels, base_channels), + DownSample(base_channels, base_channels * 2), + DownSample(base_channels * 2, base_channels * 4) + ]) + return down_samples + + def create_up_samples(self, base_channels : int) -> nn.ModuleList: + down_samples = nn.ModuleList( + [ + UpSample(base_channels * 4, base_channels * 2), + UpSample(base_channels * 2, base_channels), + UpSample(base_channels, base_channels) + ]) + return down_samples + + def forward(self, target_tensor : Tensor, target_attribute : Tensor) -> Tensor: + output_tensor = torch.cat([ target_tensor, target_attribute ], dim=1) + + for down_sample in self.down_samples: + output_tensor = down_sample(output_tensor) + output_tensor = self.bottleneck(output_tensor) + + for up_sample in self.up_samples: + output_tensor = up_sample(output_tensor) + output_tensor = self.conv(output_tensor) + output_tensor = self.activation(output_tensor) + return output_tensor + + +class ResBlock(nn.Module): + def __init__(self, channels: int): + super().__init__() + self.conv = nn.Sequential( + nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False), + nn.BatchNorm2d(channels), + nn.ReLU(inplace=True), + nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False), + nn.BatchNorm2d(channels), + nn.ReLU(inplace=True) + ) + self.relu = nn.ReLU(inplace=True) + + def forward(self, input_tensor: Tensor) -> Tensor: + output_tensor = self.conv(input_tensor) + input_tensor + output_tensor = self.relu(output_tensor) + return output_tensor + + +class UpSample(nn.Module): + def __init__(self, input_channels : int, output_channels : int) -> None: + super().__init__() + self.conv_transpose = nn.ConvTranspose2d(input_channels, output_channels, kernel_size = 2, stride = 2) + self.relu = nn.ReLU(inplace=True) + + def forward(self, input_tensor : Tensor) -> Tensor: + output_tensor = self.conv_transpose(input_tensor) + output_tensor = self.relu(output_tensor) + return output_tensor + + +class DownSample(nn.Module): + def __init__(self, input_channels : int, output_channels : int) -> None: + super().__init__() + self.conv = nn.Conv2d(input_channels, output_channels, kernel_size = 3, padding = 1, bias = False) + self.batch_norm = nn.BatchNorm2d(output_channels) + self.relu = nn.ReLU(inplace = True) + self.max_pool = nn.MaxPool2d(2) + + def forward(self, input_tensor : Tensor) -> Tensor: + output_tensor = self.conv(input_tensor) + output_tensor = self.batch_norm(output_tensor) + output_tensor = self.relu(output_tensor) + output_tensor = self.max_pool(output_tensor) + return output_tensor diff --git a/face_swapper/src/training.py b/face_swapper/src/training.py index 005753e..ada30c0 100644 --- a/face_swapper/src/training.py +++ b/face_swapper/src/training.py @@ -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 AdversarialLoss, AttributeLoss, DiscriminatorLoss, GazeLoss, IdentityLoss, MotionLoss, ReconstructionLoss +from .models.loss import AdversarialLoss, AttributeLoss, DiscriminatorLoss, GazeLoss, IdentityLoss, MaskLoss, MotionLoss, ReconstructionLoss from .types import Batch, Embedding, OptimizerSet warnings.filterwarnings('ignore', category = UserWarning, module = 'torch') @@ -31,11 +31,13 @@ class FaceSwapperTrainer(LightningModule): self.config_embedder_path = config_parser.get('training.model', 'embedder_path') self.config_gazer_path = config_parser.get('training.model', 'gazer_path') self.config_motion_extractor_path = config_parser.get('training.model', 'motion_extractor_path') + self.config_parser_path = config_parser.get('training.model', 'parser_path') self.config_learning_rate = config_parser.getfloat('training.trainer', 'learning_rate') self.config_preview_frequency = config_parser.getint('training.trainer', 'preview_frequency') self.embedder = torch.jit.load(self.config_embedder_path, map_location = 'cpu').eval() self.gazer = torch.jit.load(self.config_gazer_path, map_location = 'cpu').eval() self.motion_extractor = torch.jit.load(self.config_motion_extractor_path, map_location = 'cpu').eval() + self.parser = torch.jit.load(self.config_parser_path, map_location = 'cpu').eval() self.generator = Generator(config_parser) self.discriminator = Discriminator(config_parser) self.discriminator_loss = DiscriminatorLoss() @@ -45,11 +47,12 @@ class FaceSwapperTrainer(LightningModule): self.identity_loss = IdentityLoss(config_parser, self.embedder) self.motion_loss = MotionLoss(config_parser, self.motion_extractor) self.gaze_loss = GazeLoss(config_parser, self.gazer) + self.mask_loss = MaskLoss(config_parser, self.parser) self.automatic_optimization = False - def forward(self, source_embedding : Embedding, target_tensor : Tensor) -> Tensor: - output_tensor = self.generator(source_embedding, target_tensor) - return output_tensor + def forward(self, source_embedding : Embedding, target_tensor : Tensor) -> Tuple[Tensor, Tensor]: + output_tensor, mask_tensor = self.generator(source_embedding, target_tensor) + return output_tensor, mask_tensor def configure_optimizers(self) -> Tuple[OptimizerSet, OptimizerSet]: generator_optimizer = torch.optim.AdamW(self.generator.parameters(), lr = self.config_learning_rate, betas = (0.0, 0.999), weight_decay = 1e-4) @@ -82,7 +85,7 @@ class FaceSwapperTrainer(LightningModule): generator_optimizer, discriminator_optimizer = self.optimizers() #type:ignore[attr-defined] source_embedding = calc_embedding(self.embedder, source_tensor, (0, 0, 0, 0)) target_attributes = self.generator.get_attributes(target_tensor) - generator_output_tensor = self.generator(source_embedding, target_tensor) + generator_output_tensor, generator_mask_tensor = self.generator(source_embedding, target_tensor) generator_output_attributes = self.generator.get_attributes(generator_output_tensor) discriminator_output_tensors = self.discriminator(generator_output_tensor) @@ -93,7 +96,8 @@ class FaceSwapperTrainer(LightningModule): identity_loss, weighted_identity_loss = self.identity_loss(generator_output_tensor, source_tensor) pose_loss, weighted_pose_loss, expression_loss, weighted_expression_loss = self.motion_loss(target_tensor, generator_output_tensor) gaze_loss, weighted_gaze_loss = self.gaze_loss(target_tensor, generator_output_tensor) - generator_loss = weighted_adversarial_loss + weighted_attribute_loss + weighted_reconstruction_loss + weighted_identity_loss + weighted_pose_loss + weighted_gaze_loss + weighted_expression_loss + mask_loss, weighted_mask_loss = self.mask_loss(target_tensor, generator_mask_tensor) + generator_loss = weighted_adversarial_loss + weighted_attribute_loss + weighted_reconstruction_loss + weighted_identity_loss + weighted_pose_loss + weighted_gaze_loss + weighted_expression_loss + weighted_mask_loss generator_optimizer.zero_grad() self.manual_backward(generator_loss) @@ -121,12 +125,13 @@ class FaceSwapperTrainer(LightningModule): self.log('identity_loss', identity_loss) self.log('pose_loss', pose_loss) self.log('gaze_loss', gaze_loss) + self.log('mask_loss', mask_loss) 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)) - output_tensor = self.generator(source_embedding, target_tensor) + output_tensor, mask_tensor = self.generator(source_embedding, target_tensor) output_embedding = calc_embedding(self.embedder, output_tensor, (0, 0, 0, 0)) validation_score = (nn.functional.cosine_similarity(source_embedding, output_embedding).mean() + 1) * 0.5 self.log('validation_score', validation_score, prog_bar = True) diff --git a/face_swapper/src/types.py b/face_swapper/src/types.py index 9f0c526..bce6a9e 100644 --- a/face_swapper/src/types.py +++ b/face_swapper/src/types.py @@ -16,6 +16,7 @@ GeneratorModule : TypeAlias = Module EmbedderModule : TypeAlias = Module GazerModule : TypeAlias = Module MotionExtractorModule : TypeAlias = Module +ParserModule : TypeAlias = Module OptimizerSet : TypeAlias = Any