mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-07-28 16:08:50 +02:00
Rework on config
This commit is contained in:
@@ -1,31 +1,28 @@
|
||||
import configparser
|
||||
from configparser import ConfigParser
|
||||
from typing import List
|
||||
|
||||
from torch import Tensor, nn
|
||||
|
||||
from ..networks.nld import NLD
|
||||
|
||||
CONFIG = configparser.ConfigParser()
|
||||
CONFIG.read('config.ini')
|
||||
|
||||
|
||||
class Discriminator(nn.Module):
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, config_parser : ConfigParser) -> None:
|
||||
super().__init__()
|
||||
self.config =\
|
||||
{
|
||||
'num_discriminators': config_parser.getint('training.model.discriminator', 'num_discriminators')
|
||||
}
|
||||
self.config_parser = config_parser
|
||||
self.avg_pool = nn.AvgPool2d(kernel_size = 3, stride = 2, padding = (1, 1), count_include_pad = False)
|
||||
self.discriminators = self.create_discriminators()
|
||||
|
||||
@staticmethod
|
||||
def create_discriminators() -> nn.ModuleList:
|
||||
num_discriminators = CONFIG.getint('training.model.discriminator', 'num_discriminators')
|
||||
input_channels = CONFIG.getint('training.model.discriminator', 'input_channels')
|
||||
num_filters = CONFIG.getint('training.model.discriminator', 'num_filters')
|
||||
kernel_size = CONFIG.getint('training.model.discriminator', 'kernel_size')
|
||||
num_layers = CONFIG.getint('training.model.discriminator', 'num_layers')
|
||||
|
||||
def create_discriminators(self) -> nn.ModuleList:
|
||||
discriminators = nn.ModuleList()
|
||||
|
||||
for _ in range(num_discriminators):
|
||||
discriminator = NLD(input_channels, num_filters, num_layers, kernel_size).sequences
|
||||
for _ in range(self.config.get('num_discriminators')):
|
||||
discriminator = NLD(self.config_parser).sequences
|
||||
discriminators.append(discriminator)
|
||||
|
||||
return discriminators
|
||||
@@ -35,7 +32,8 @@ class Discriminator(nn.Module):
|
||||
output_tensors = []
|
||||
|
||||
for discriminator in self.discriminators:
|
||||
output_tensors.append(discriminator(temp_tensor))
|
||||
output_tensor = discriminator(temp_tensor)
|
||||
output_tensors.append(output_tensor)
|
||||
temp_tensor = self.avg_pool(temp_tensor)
|
||||
|
||||
return output_tensors
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import configparser
|
||||
from configparser import ConfigParser
|
||||
|
||||
from torch import Tensor, nn
|
||||
|
||||
@@ -6,20 +6,12 @@ from ..networks.aad import AAD
|
||||
from ..networks.unet import UNet
|
||||
from ..types import Attributes, Embedding
|
||||
|
||||
CONFIG = configparser.ConfigParser()
|
||||
CONFIG.read('config.ini')
|
||||
|
||||
|
||||
class Generator(nn.Module):
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, config_parser : ConfigParser) -> None:
|
||||
super().__init__()
|
||||
identity_channels = CONFIG.getint('training.model.generator', 'identity_channels')
|
||||
output_channels = CONFIG.getint('training.model.generator', 'output_channels')
|
||||
output_size = CONFIG.getint('training.model.generator', 'output_size')
|
||||
num_blocks = CONFIG.getint('training.model.generator', 'num_blocks')
|
||||
|
||||
self.encoder = UNet(output_size)
|
||||
self.generator = AAD(identity_channels, output_channels, output_size, num_blocks)
|
||||
self.encoder = UNet(config_parser)
|
||||
self.generator = AAD(config_parser)
|
||||
self.encoder.apply(init_weight)
|
||||
self.generator.apply(init_weight)
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import configparser
|
||||
from configparser import ConfigParser
|
||||
from typing import List, Tuple
|
||||
|
||||
import torch
|
||||
@@ -9,9 +9,6 @@ from torchvision import transforms
|
||||
from ..helper import calc_embedding
|
||||
from ..types import Attributes, EmbedderModule, Gaze, GazerModule, MotionExtractorModule
|
||||
|
||||
CONFIG = configparser.ConfigParser()
|
||||
CONFIG.read('config.ini')
|
||||
|
||||
|
||||
class DiscriminatorLoss(nn.Module):
|
||||
def __init__(self) -> None:
|
||||
@@ -36,11 +33,14 @@ class DiscriminatorLoss(nn.Module):
|
||||
|
||||
|
||||
class AdversarialLoss(nn.Module):
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, config_parser : ConfigParser) -> None:
|
||||
super().__init__()
|
||||
self.config =\
|
||||
{
|
||||
'adversarial_weight': config_parser.getfloat('training.losses', 'adversarial_weight')
|
||||
}
|
||||
|
||||
def forward(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:
|
||||
@@ -48,7 +48,7 @@ class AdversarialLoss(nn.Module):
|
||||
temp_tensors.append(temp_tensor)
|
||||
|
||||
adversarial_loss = torch.stack(temp_tensors).mean()
|
||||
weighted_adversarial_loss = adversarial_loss * adversarial_weight
|
||||
weighted_adversarial_loss = adversarial_loss * self.config.get('adversarial_weight')
|
||||
return adversarial_loss, weighted_adversarial_loss
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user