mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-06-25 07:59:55 +02:00
Improve optimizer configs
This commit is contained in:
@@ -53,8 +53,7 @@ class EmbeddingConverterTrainer(lightning.LightningModule):
|
||||
learning_rate = CONFIG.getfloat('training.trainer', 'learning_rate')
|
||||
optimizer = torch.optim.Adam(self.parameters(), lr = learning_rate)
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer)
|
||||
|
||||
return\
|
||||
config =\
|
||||
{
|
||||
'optimizer': optimizer,
|
||||
'lr_scheduler':
|
||||
@@ -66,6 +65,8 @@ class EmbeddingConverterTrainer(lightning.LightningModule):
|
||||
}
|
||||
}
|
||||
|
||||
return config
|
||||
|
||||
|
||||
def create_loaders(dataset : Dataset[Tensor]) -> Tuple[DataLoader[Tensor], DataLoader[Tensor]]:
|
||||
batch_size = CONFIG.getint('training.loader', 'batch_size')
|
||||
|
||||
@@ -17,7 +17,7 @@ from .helper import calc_embedding
|
||||
from .models.discriminator import Discriminator
|
||||
from .models.generator import Generator
|
||||
from .models.loss import AdversarialLoss, AttributeLoss, DiscriminatorLoss, GazeLoss, IdentityLoss, PoseLoss, ReconstructionLoss
|
||||
from .types import Batch, Embedding
|
||||
from .types import Batch, Embedding, OptimizerConfig
|
||||
|
||||
CONFIG = configparser.ConfigParser()
|
||||
CONFIG.read('config.ini')
|
||||
@@ -44,11 +44,36 @@ class FaceSwapperTrainer(lightning.LightningModule):
|
||||
output_tensor = self.generator(source_embedding, target_tensor)
|
||||
return output_tensor
|
||||
|
||||
def configure_optimizers(self) -> Tuple[Optimizer, Optimizer]:
|
||||
def configure_optimizers(self) -> Tuple[OptimizerConfig, OptimizerConfig]:
|
||||
learning_rate = CONFIG.getfloat('training.trainer', 'learning_rate')
|
||||
generator_optimizer = torch.optim.Adam(self.generator.parameters(), lr = learning_rate, betas = (0.0, 0.999), weight_decay = 1e-4)
|
||||
discriminator_optimizer = torch.optim.Adam(self.discriminator.parameters(), lr = learning_rate, betas = (0.0, 0.999), weight_decay = 1e-4)
|
||||
return generator_optimizer, discriminator_optimizer
|
||||
generator_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(generator_optimizer)
|
||||
discriminator_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(discriminator_optimizer)
|
||||
|
||||
generator_config =\
|
||||
{
|
||||
'optimizer': generator_optimizer,
|
||||
'lr_scheduler':
|
||||
{
|
||||
'scheduler': generator_scheduler,
|
||||
'monitor': 'generator_loss',
|
||||
'interval': 'step',
|
||||
'frequency': 1000
|
||||
}
|
||||
}
|
||||
discriminator_config =\
|
||||
{
|
||||
'optimizer': discriminator_optimizer,
|
||||
'lr_scheduler':
|
||||
{
|
||||
'scheduler': discriminator_scheduler,
|
||||
'monitor': 'discriminator_loss',
|
||||
'interval': 'step',
|
||||
'frequency': 1000
|
||||
}
|
||||
}
|
||||
return generator_config, discriminator_config
|
||||
|
||||
def training_step(self, batch : Batch, batch_index : int) -> Tensor:
|
||||
preview_frequency = CONFIG.getfloat('training.trainer', 'preview_frequency')
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Tuple, TypeAlias
|
||||
from typing import Any, Tuple, TypeAlias
|
||||
|
||||
from torch import Tensor
|
||||
from torch.nn import Module
|
||||
@@ -13,3 +13,5 @@ Padding : TypeAlias = Tuple[int, int, int, int]
|
||||
|
||||
GeneratorModule : TypeAlias = Module
|
||||
EmbedderModule : TypeAlias = Module
|
||||
|
||||
OptimizerConfig : TypeAlias = Any
|
||||
|
||||
Reference in New Issue
Block a user