diff --git a/embedding_converter/src/training.py b/embedding_converter/src/training.py index fedeffa..331ce39 100644 --- a/embedding_converter/src/training.py +++ b/embedding_converter/src/training.py @@ -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') diff --git a/face_swapper/src/training.py b/face_swapper/src/training.py index 56cacb4..dfbf9c6 100644 --- a/face_swapper/src/training.py +++ b/face_swapper/src/training.py @@ -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') diff --git a/face_swapper/src/types.py b/face_swapper/src/types.py index 27058ed..cc2f46e 100644 --- a/face_swapper/src/types.py +++ b/face_swapper/src/types.py @@ -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