mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-06-25 07:59:55 +02:00
add masknet
This commit is contained in:
@@ -74,6 +74,7 @@ identity_weight = 20.0
|
||||
gaze_weight = 0.0
|
||||
pose_weight = 0.0
|
||||
expression_weight = 0.0
|
||||
mask_weight = 1.0
|
||||
```
|
||||
|
||||
```
|
||||
|
||||
@@ -36,6 +36,7 @@ identity_weight =
|
||||
gaze_weight =
|
||||
pose_weight =
|
||||
expression_weight =
|
||||
mask_weight =
|
||||
|
||||
[training.trainer]
|
||||
learning_rate =
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -16,6 +16,7 @@ GeneratorModule : TypeAlias = Module
|
||||
EmbedderModule : TypeAlias = Module
|
||||
GazerModule : TypeAlias = Module
|
||||
MotionExtractorModule : TypeAlias = Module
|
||||
ParserModule : TypeAlias = Module
|
||||
|
||||
OptimizerSet : TypeAlias = Any
|
||||
|
||||
|
||||
Reference in New Issue
Block a user