Remove annoying Tuner, Match both trainer configs a bit more

This commit is contained in:
henryruhs
2025-02-24 15:40:50 +01:00
parent bfcbd6bf95
commit 5b3b2abdd7
6 changed files with 7 additions and 10 deletions
+1
View File
@@ -47,6 +47,7 @@ target_path = .models/arcface_simswap.pt
[training.trainer]
learning_rate = 0.001
max_epochs = 4096
precision = 16-mixed
```
```
+1
View File
@@ -13,6 +13,7 @@ target_path =
[training.trainer]
learning_rate =
max_epochs =
precision =
[training.output]
directory_path =
+4 -6
View File
@@ -7,7 +7,6 @@ import torch
from lightning import Trainer
from lightning.pytorch.callbacks import ModelCheckpoint
from lightning.pytorch.loggers import TensorBoardLogger
from lightning.pytorch.tuner import Tuner
from torch import Tensor, nn
from torch.utils.data import DataLoader, Dataset, random_split
@@ -24,13 +23,11 @@ class EmbeddingConverterTrainer(lightning.LightningModule):
super(EmbeddingConverterTrainer, self).__init__()
source_path = CONFIG.get('training.model', 'source_path')
target_path = CONFIG.get('training.model', 'target_path')
learning_rate = CONFIG.getfloat('training.trainer', 'learning_rate')
self.embedding_converter = EmbeddingConverter()
self.source_embedder = torch.jit.load(source_path, map_location = 'cpu') # type:ignore[no-untyped-call]
self.target_embedder = torch.jit.load(target_path, map_location = 'cpu') # type:ignore[no-untyped-call]
self.mse_loss = nn.MSELoss()
self.lr = learning_rate
def forward(self, source_embedding : Embedding) -> Embedding:
return self.embedding_converter(source_embedding)
@@ -93,19 +90,22 @@ def create_trainer() -> Trainer:
trainer_max_epochs = CONFIG.getint('training.trainer', 'max_epochs')
output_directory_path = CONFIG.get('training.output', 'directory_path')
output_file_pattern = CONFIG.get('training.output', 'file_pattern')
trainer_precision = CONFIG.get('training.trainer', 'precision')
logger = TensorBoardLogger('.logs', name = 'embedding_converter')
os.makedirs(output_directory_path, exist_ok = True)
return Trainer(
logger = logger,
log_every_n_steps = 10,
max_epochs = trainer_max_epochs,
precision = trainer_precision, # type:ignore[arg-type]
callbacks =
[
ModelCheckpoint(
monitor = 'training_loss',
dirpath = output_directory_path,
filename = output_file_pattern,
every_n_epochs = 10,
every_n_epochs = 1,
save_top_k = 3,
save_last = True
)
@@ -121,8 +121,6 @@ def train() -> None:
training_loader, validation_loader = create_loaders(dataset)
embedding_converter_trainer = EmbeddingConverterTrainer()
trainer = create_trainer()
tuner = Tuner(trainer)
tuner.lr_find(embedding_converter_trainer, training_loader, validation_loader)
if os.path.exists(output_resume_path):
trainer.fit(embedding_converter_trainer, training_loader, validation_loader, ckpt_path = output_resume_path)
-1
View File
@@ -77,7 +77,6 @@ gaze_weight = 0
learning_rate = 0.0004
max_epochs = 50
precision = 16-mixed
automatic_optimization = false
preview_frequency = 250
```
-1
View File
@@ -37,7 +37,6 @@ gaze_weight =
learning_rate =
max_epochs =
precision =
automatic_optimization =
preview_frequency =
[training.output]
+1 -2
View File
@@ -26,7 +26,6 @@ CONFIG.read('config.ini')
class FaceSwapperTrainer(lightning.LightningModule):
def __init__(self) -> None:
super().__init__()
automatic_optimization = CONFIG.getboolean('training.trainer', 'automatic_optimization')
embedder_path = CONFIG.get('training.model', 'embedder_path')
self.generator = Generator()
@@ -39,7 +38,7 @@ class FaceSwapperTrainer(lightning.LightningModule):
self.pose_loss = PoseLoss()
self.gaze_loss = GazeLoss()
self.embedder = torch.jit.load(embedder_path, map_location = 'cpu') # type:ignore[no-untyped-call]
self.automatic_optimization = automatic_optimization
self.automatic_optimization = False
def forward(self, target_tensor : Tensor, source_embedding : Embedding) -> Tensor:
output_tensor = self.generator(source_embedding, target_tensor)