Improve optimizer configs

This commit is contained in:
henryruhs
2025-02-26 00:01:18 +01:00
parent 484a49c27d
commit 0ad2556c4c
3 changed files with 34 additions and 6 deletions
+3 -2
View File
@@ -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')
+28 -3
View File
@@ -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')
+3 -1
View File
@@ -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