add masknet

This commit is contained in:
harisreedhar
2025-03-08 16:05:07 +05:30
committed by henryruhs
parent 4af22832db
commit 8f1f002c64
7 changed files with 142 additions and 10 deletions
+1
View File
@@ -74,6 +74,7 @@ identity_weight = 20.0
gaze_weight = 0.0
pose_weight = 0.0
expression_weight = 0.0
mask_weight = 1.0
```
```
+1
View File
@@ -36,6 +36,7 @@ identity_weight =
gaze_weight =
pose_weight =
expression_weight =
mask_weight =
[training.trainer]
learning_rate =
+7 -2
View File
@@ -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)
+30 -1
View File
@@ -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
+90
View File
@@ -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
+12 -7
View File
@@ -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)
+1
View File
@@ -16,6 +16,7 @@ GeneratorModule : TypeAlias = Module
EmbedderModule : TypeAlias = Module
GazerModule : TypeAlias = Module
MotionExtractorModule : TypeAlias = Module
ParserModule : TypeAlias = Module
OptimizerSet : TypeAlias = Any