mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-06-25 07:59:55 +02:00
changes
This commit is contained in:
@@ -3,15 +3,15 @@ import os
|
||||
from typing import Any, Tuple
|
||||
|
||||
import lightning
|
||||
import numpy
|
||||
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, TensorDataset, random_split
|
||||
from torch.utils.data import DataLoader, Dataset, random_split
|
||||
|
||||
from .data_loader import DataLoaderRecognition
|
||||
from .models.embedding_converter import EmbeddingConverter
|
||||
from .types import Batch, Embedding
|
||||
|
||||
@@ -22,23 +22,34 @@ CONFIG.read('config.ini')
|
||||
class EmbeddingConverterTrainer(lightning.LightningModule):
|
||||
def __init__(self) -> None:
|
||||
super(EmbeddingConverterTrainer, self).__init__()
|
||||
source_path = CONFIG.get('training.model', 'source_path')
|
||||
target_path = CONFIG.get('training.model', 'target_path')
|
||||
self.lr = CONFIG.getfloat('training.trainer', 'learning_rate')
|
||||
self.embedding_converter = EmbeddingConverter()
|
||||
self.source_embedder = torch.jit.load(source_path) # type:ignore[no-untyped-call]
|
||||
self.target_embedder = torch.jit.load(target_path) # type:ignore[no-untyped-call]
|
||||
self.source_embedder.eval()
|
||||
self.target_embedder.eval()
|
||||
self.mse_loss = nn.MSELoss()
|
||||
|
||||
def forward(self, source_embedding : Embedding) -> Embedding:
|
||||
return self.embedding_converter(source_embedding)
|
||||
|
||||
def training_step(self, batch : Batch, batch_index : int) -> Tensor:
|
||||
source_tensor, target = batch
|
||||
output_tensor = self(source_tensor)
|
||||
loss_training = self.mse_loss(output_tensor, target)
|
||||
with torch.no_grad():
|
||||
source_embedding = self.source_embedder(batch)
|
||||
target_embedding = self.target_embedder(batch)
|
||||
output_embedding = self(source_embedding)
|
||||
loss_training = self.mse_loss(output_embedding, target_embedding)
|
||||
self.log('loss_training', loss_training, prog_bar = True)
|
||||
return loss_training
|
||||
|
||||
def validation_step(self, batch : Batch, batch_index : int) -> Tensor:
|
||||
source_tensor, target_tensor = batch
|
||||
output_tensor = self(source_tensor)
|
||||
validation = self.mse_loss(output_tensor, target_tensor)
|
||||
with torch.no_grad():
|
||||
source_embedding = self.source_embedder(batch)
|
||||
target_embedding = self.target_embedder(batch)
|
||||
output_embedding = self(source_embedding)
|
||||
validation = self.mse_loss(output_embedding, target_embedding)
|
||||
self.log('validation', validation, prog_bar = True)
|
||||
return validation
|
||||
|
||||
@@ -53,36 +64,28 @@ class EmbeddingConverterTrainer(lightning.LightningModule):
|
||||
'lr_scheduler':
|
||||
{
|
||||
'scheduler': scheduler,
|
||||
'monitor': 'train_loss',
|
||||
'monitor': 'loss_training',
|
||||
'interval': 'epoch',
|
||||
'frequency': 1
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def create_loaders() -> Tuple[DataLoader, DataLoader]:
|
||||
loader_batch_size = CONFIG.getint('training.loader', 'batch_size')
|
||||
loader_num_workers = CONFIG.getint('training.loader', 'num_workers')
|
||||
def create_loaders(dataset : Dataset[Any]) -> Tuple[DataLoader[Any], DataLoader[Any]]:
|
||||
batch_size = CONFIG.getint('training.loader', 'batch_size')
|
||||
num_workers = CONFIG.getint('training.loader', 'num_workers')
|
||||
|
||||
training_dataset, validate_dataset = split_dataset()
|
||||
training_loader = DataLoader(training_dataset, batch_size = loader_batch_size, num_workers = loader_num_workers, shuffle = True, pin_memory = True)
|
||||
validation_loader = DataLoader(validate_dataset, batch_size = loader_batch_size, num_workers = loader_num_workers, shuffle = False, pin_memory = True)
|
||||
training_dataset, validate_dataset = split_dataset(dataset)
|
||||
training_loader = DataLoader(training_dataset, batch_size = batch_size, shuffle = True, num_workers = num_workers, drop_last = True, pin_memory = True, persistent_workers = True)
|
||||
validation_loader = DataLoader(validate_dataset, batch_size = batch_size, shuffle = False, num_workers = num_workers, drop_last = False, pin_memory = True, persistent_workers = True)
|
||||
return training_loader, validation_loader
|
||||
|
||||
|
||||
def split_dataset() -> Tuple[Dataset[Any], Dataset[Any]]:
|
||||
input_source_path = CONFIG.get('preparing.input', 'source_path')
|
||||
input_target_path = CONFIG.get('preparing.input', 'target_path')
|
||||
def split_dataset(dataset : Dataset[Any]) -> Tuple[Dataset[Any], Dataset[Any]]:
|
||||
loader_split_ratio = CONFIG.getfloat('training.loader', 'split_ratio')
|
||||
|
||||
source_tensor = torch.from_numpy(numpy.load(input_source_path)).float()
|
||||
target_tensor = torch.from_numpy(numpy.load(input_target_path)).float()
|
||||
dataset = TensorDataset(source_tensor, target_tensor)
|
||||
|
||||
dataset_size = len(dataset)
|
||||
training_size = int(loader_split_ratio * len(dataset))
|
||||
validation_size = int(dataset_size - training_size)
|
||||
training_dataset, validate_dataset = random_split(dataset, [ training_size, validation_size ])
|
||||
training_size = int(loader_split_ratio * len(dataset)) # type:ignore[operator, arg-type]
|
||||
validation_size = len(dataset) - training_size # type:ignore[arg-type]
|
||||
training_dataset, validate_dataset = random_split(dataset, [training_size, validation_size])
|
||||
return training_dataset, validate_dataset
|
||||
|
||||
|
||||
@@ -112,9 +115,12 @@ def create_trainer() -> Trainer:
|
||||
|
||||
|
||||
def train() -> None:
|
||||
dataset_path = CONFIG.get('preparing.dataset', 'dataset_path')
|
||||
dataset_image_pattern = CONFIG.get('preparing.dataset', 'image_pattern')
|
||||
resume_file_path = CONFIG.get('training.output', 'resume_file_path')
|
||||
|
||||
training_loader, validation_loader = create_loaders()
|
||||
dataset = DataLoaderRecognition(dataset_path, dataset_image_pattern)
|
||||
training_loader, validation_loader = create_loaders(dataset)
|
||||
embedding_converter_trainer = EmbeddingConverterTrainer()
|
||||
trainer = create_trainer()
|
||||
tuner = Tuner(trainer)
|
||||
|
||||
Reference in New Issue
Block a user