mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-07-28 16:08:50 +02:00
Join MaskNet to guide generator
This commit is contained in:
@@ -5,7 +5,8 @@ from torch import Tensor, nn
|
||||
|
||||
from ..networks.aad import AAD
|
||||
from ..networks.unet import UNet
|
||||
from ..types import Embedding, Feature
|
||||
from ..networks.masknet import MaskNet
|
||||
from ..types import Embedding, Feature, Mask
|
||||
|
||||
|
||||
class Generator(nn.Module):
|
||||
@@ -13,13 +14,17 @@ class Generator(nn.Module):
|
||||
super().__init__()
|
||||
self.encoder = UNet(config_parser)
|
||||
self.generator = AAD(config_parser)
|
||||
self.masker = MaskNet(config_parser)
|
||||
self.encoder.apply(init_weight)
|
||||
self.generator.apply(init_weight)
|
||||
self.masker.apply(init_weight)
|
||||
|
||||
def forward(self, source_embedding : Embedding, target_tensor : Tensor) -> Tuple[Tensor, Tuple[Feature, ...]]:
|
||||
def forward(self, source_embedding : Embedding, target_tensor : Tensor) -> Tuple[Tensor, Mask]:
|
||||
target_features = self.encode_features(target_tensor)
|
||||
output_tensor = self.generator(source_embedding, target_features)
|
||||
return output_tensor, target_features
|
||||
target_feature = target_features[-1]
|
||||
output_mask = self.masker(target_tensor, target_feature)
|
||||
return output_tensor, output_mask
|
||||
|
||||
def encode_features(self, input_tensor : Tensor) -> Tuple[Feature, ...]:
|
||||
return self.encoder(input_tensor)
|
||||
|
||||
@@ -182,16 +182,18 @@ class GazeLoss(nn.Module):
|
||||
class MaskLoss(nn.Module):
|
||||
def __init__(self, config_parser : ConfigParser, face_parser : FaceParserModule) -> 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.face_parser = face_parser
|
||||
self.mse_loss = nn.MSELoss()
|
||||
|
||||
def forward(self, target_tensor : Tensor, output_mask : Mask) -> Loss:
|
||||
def forward(self, target_tensor : Tensor, output_mask : Mask) -> Tuple[Loss, Loss]:
|
||||
target_mask = self.calc_mask(target_tensor)
|
||||
target_mask = target_mask.view(-1, self.config_output_size, self.config_output_size)
|
||||
output_mask = output_mask.view(-1, self.config_output_size, self.config_output_size)
|
||||
mask_loss = self.mse_loss(target_mask, output_mask)
|
||||
return mask_loss
|
||||
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')
|
||||
|
||||
Reference in New Issue
Block a user