mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-06-25 07:59:55 +02:00
Final rename for everything
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
OpenRAIL-MS license
|
||||
|
||||
Copyright (c) 2025 Henry Ruhs
|
||||
@@ -0,0 +1,96 @@
|
||||
CrossFace
|
||||
=========
|
||||
|
||||
> Seamless transform face embeddings across embedder models.
|
||||
|
||||

|
||||
|
||||
|
||||
Preview
|
||||
-------
|
||||
|
||||

|
||||
|
||||
|
||||
Installation
|
||||
------------
|
||||
|
||||
```
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
|
||||
Setup
|
||||
-----
|
||||
|
||||
This `config.ini` utilizes the MegaFace dataset to train the CrossFace model for SimSwap.
|
||||
|
||||
```
|
||||
[training.dataset]
|
||||
file_pattern = .datasets/megaface/**/*.jpg
|
||||
```
|
||||
|
||||
```
|
||||
[training.loader]
|
||||
batch_size = 256
|
||||
num_workers = 8
|
||||
split_ratio = 0.95
|
||||
```
|
||||
|
||||
```
|
||||
[training.model]
|
||||
source_path = .models/arcface_w600k_r50.pt
|
||||
target_path = .models/arcface_simswap.pt
|
||||
```
|
||||
|
||||
```
|
||||
[training.trainer]
|
||||
learning_rate = 0.001
|
||||
max_epochs = 4096
|
||||
strategy = auto
|
||||
precision = 16-mixed
|
||||
logger_path = .logs
|
||||
logger_name = crossface_simswap
|
||||
```
|
||||
|
||||
```
|
||||
[training.output]
|
||||
directory_path = .outputs
|
||||
file_pattern = crossface_simswap_{epoch}_{step}
|
||||
resume_path = .outputs/last.ckpt
|
||||
```
|
||||
|
||||
```
|
||||
[exporting]
|
||||
directory_path = .exports
|
||||
source_path = .outputs/last.ckpt
|
||||
target_path = .exports/crossface_simswap.onnx
|
||||
ir_version = 10
|
||||
opset_version = 15
|
||||
```
|
||||
|
||||
|
||||
Training
|
||||
--------
|
||||
|
||||
Train the model.
|
||||
|
||||
```
|
||||
python train.py
|
||||
```
|
||||
|
||||
Launch the TensorBoard to monitor the training.
|
||||
|
||||
```
|
||||
tensorboard --logdir=.logs
|
||||
```
|
||||
|
||||
|
||||
Exporting
|
||||
---------
|
||||
|
||||
Export the model to ONNX.
|
||||
|
||||
```
|
||||
python export.py
|
||||
```
|
||||
@@ -0,0 +1,31 @@
|
||||
[training.dataset]
|
||||
file_pattern =
|
||||
|
||||
[training.loader]
|
||||
batch_size =
|
||||
num_workers =
|
||||
split_ratio =
|
||||
|
||||
[training.model]
|
||||
source_path =
|
||||
target_path =
|
||||
|
||||
[training.trainer]
|
||||
learning_rate =
|
||||
max_epochs =
|
||||
strategy =
|
||||
precision =
|
||||
logger_path =
|
||||
logger_name =
|
||||
|
||||
[training.output]
|
||||
directory_path =
|
||||
file_pattern =
|
||||
resume_path =
|
||||
|
||||
[exporting]
|
||||
directory_path =
|
||||
source_path =
|
||||
target_path =
|
||||
ir_version =
|
||||
opset_version =
|
||||
@@ -0,0 +1,6 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
from src.exporting import export
|
||||
|
||||
if __name__ == '__main__':
|
||||
export()
|
||||
@@ -0,0 +1,34 @@
|
||||
import glob
|
||||
from configparser import ConfigParser
|
||||
|
||||
from torch import Tensor
|
||||
from torch.utils.data import Dataset
|
||||
from torchvision import io, transforms
|
||||
|
||||
from .types import Batch
|
||||
|
||||
|
||||
class StaticDataset(Dataset[Tensor]):
|
||||
def __init__(self, config_parser : ConfigParser) -> None:
|
||||
self.config_file_pattern = config_parser.get('training.dataset', 'file_pattern')
|
||||
self.file_paths = glob.glob(self.config_file_pattern)
|
||||
self.transforms = self.compose_transforms()
|
||||
|
||||
def __getitem__(self, index : int) -> Batch:
|
||||
file_path = self.file_paths[index]
|
||||
temp_tensor = io.read_image(file_path)
|
||||
return self.transforms(temp_tensor)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.file_paths)
|
||||
|
||||
@staticmethod
|
||||
def compose_transforms() -> transforms:
|
||||
return transforms.Compose(
|
||||
[
|
||||
transforms.ToPILImage(),
|
||||
transforms.Resize((112, 112), interpolation = transforms.InterpolationMode.BICUBIC),
|
||||
transforms.ColorJitter(brightness = 0.2, contrast = 0.2, saturation = 0.2, hue = 0.1),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
|
||||
])
|
||||
@@ -0,0 +1,23 @@
|
||||
import os
|
||||
from configparser import ConfigParser
|
||||
|
||||
import torch
|
||||
|
||||
from .training import CrossFaceTrainer
|
||||
|
||||
CONFIG_PARSER = ConfigParser()
|
||||
CONFIG_PARSER.read('config.ini')
|
||||
|
||||
|
||||
def export() -> None:
|
||||
config_directory_path = CONFIG_PARSER.get('exporting', 'directory_path')
|
||||
config_source_path = CONFIG_PARSER.get('exporting', 'source_path')
|
||||
config_target_path = CONFIG_PARSER.get('exporting', 'target_path')
|
||||
config_ir_version = CONFIG_PARSER.getint('exporting', 'ir_version')
|
||||
config_opset_version = CONFIG_PARSER.getint('exporting', 'opset_version')
|
||||
|
||||
os.makedirs(config_directory_path, exist_ok = True)
|
||||
model = CrossFaceTrainer.load_from_checkpoint(config_source_path, map_location ='cpu').eval()
|
||||
model.ir_version = torch.tensor(config_ir_version)
|
||||
input_tensor = torch.randn(1, 512)
|
||||
torch.onnx.export(model, input_tensor, config_target_path, input_names = [ 'input' ], output_names = [ 'output' ], opset_version = config_opset_version)
|
||||
@@ -0,0 +1,28 @@
|
||||
import torch
|
||||
from torch import Tensor, nn
|
||||
|
||||
|
||||
class CrossFace(nn.Module):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.layers = self.create_layers()
|
||||
self.leaky_relu = nn.LeakyReLU()
|
||||
|
||||
@staticmethod
|
||||
def create_layers() -> nn.ModuleList:
|
||||
return nn.ModuleList(
|
||||
[
|
||||
nn.Linear(512, 1024),
|
||||
nn.Linear(1024, 2048),
|
||||
nn.Linear(2048, 1024),
|
||||
nn.Linear(1024, 512)
|
||||
])
|
||||
|
||||
def forward(self, input_tensor : Tensor) -> Tensor:
|
||||
output_tensor = input_tensor / torch.norm(input_tensor)
|
||||
|
||||
for layer in self.layers[:-1]:
|
||||
output_tensor = self.leaky_relu(layer(output_tensor))
|
||||
|
||||
output_tensor = self.layers[-1](output_tensor)
|
||||
return output_tensor
|
||||
@@ -0,0 +1,134 @@
|
||||
import os
|
||||
from configparser import ConfigParser
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from lightning import LightningModule, Trainer
|
||||
from lightning.pytorch.callbacks import ModelCheckpoint
|
||||
from lightning.pytorch.loggers import TensorBoardLogger
|
||||
from torch import Tensor, nn
|
||||
from torch.utils.data import Dataset, random_split
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from .dataset import StaticDataset
|
||||
from .models.crossface import CrossFace
|
||||
from .types import Batch, Embedding, OptimizerSet
|
||||
|
||||
CONFIG_PARSER = ConfigParser()
|
||||
CONFIG_PARSER.read('config.ini')
|
||||
|
||||
|
||||
class CrossFaceTrainer(LightningModule):
|
||||
def __init__(self, config_parser : ConfigParser) -> None:
|
||||
super().__init__()
|
||||
self.config_source_path = config_parser.get('training.model', 'source_path')
|
||||
self.config_target_path = config_parser.get('training.model', 'target_path')
|
||||
self.config_learning_rate = config_parser.getfloat('training.trainer', 'learning_rate')
|
||||
self.crossface = CrossFace()
|
||||
self.source_embedder = torch.jit.load(self.config_source_path, map_location = 'cpu').eval()
|
||||
self.target_embedder = torch.jit.load(self.config_target_path, map_location = 'cpu').eval()
|
||||
self.mse_loss = nn.MSELoss()
|
||||
|
||||
def forward(self, source_embedding : Embedding) -> Embedding:
|
||||
return self.crossface(source_embedding)
|
||||
|
||||
def training_step(self, batch : Batch, batch_index : int) -> Tensor:
|
||||
with torch.no_grad():
|
||||
source_embedding = self.source_embedder(batch)
|
||||
target_embedding = self.target_embedder(batch)
|
||||
output_embedding = self(source_embedding)
|
||||
training_loss = self.mse_loss(output_embedding, target_embedding)
|
||||
self.log('training_loss', training_loss, prog_bar = True)
|
||||
return training_loss
|
||||
|
||||
def validation_step(self, batch : Batch, batch_index : int) -> Tensor:
|
||||
with torch.no_grad():
|
||||
source_embedding = self.source_embedder(batch)
|
||||
output_embedding = self(source_embedding)
|
||||
validation_score = (nn.functional.cosine_similarity(source_embedding, output_embedding).mean() + 1) * 0.5
|
||||
self.log('validation_score', validation_score, sync_dist = True, prog_bar = True)
|
||||
return validation_score
|
||||
|
||||
def configure_optimizers(self) -> OptimizerSet:
|
||||
optimizer = torch.optim.Adam(self.parameters(), lr = self.config_learning_rate)
|
||||
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer)
|
||||
optimizer_set =\
|
||||
{
|
||||
'optimizer': optimizer,
|
||||
'lr_scheduler':
|
||||
{
|
||||
'scheduler': scheduler,
|
||||
'monitor': 'training_loss',
|
||||
'interval': 'epoch',
|
||||
'frequency': 1
|
||||
}
|
||||
}
|
||||
|
||||
return optimizer_set
|
||||
|
||||
|
||||
def create_loaders(dataset : Dataset[Tensor]) -> Tuple[StatefulDataLoader[Tensor], StatefulDataLoader[Tensor]]:
|
||||
config_batch_size = CONFIG_PARSER.getint('training.loader', 'batch_size')
|
||||
config_num_workers = CONFIG_PARSER.getint('training.loader', 'num_workers')
|
||||
|
||||
training_dataset, validate_dataset = split_dataset(dataset)
|
||||
training_loader = StatefulDataLoader(training_dataset, batch_size = config_batch_size, shuffle = True, num_workers = config_num_workers, drop_last = True, pin_memory = True, persistent_workers = True)
|
||||
validation_loader = StatefulDataLoader(validate_dataset, batch_size = config_batch_size, shuffle = False, num_workers = config_num_workers, pin_memory = True, persistent_workers = True)
|
||||
return training_loader, validation_loader
|
||||
|
||||
|
||||
def split_dataset(dataset : Dataset[Tensor]) -> Tuple[Dataset[Tensor], Dataset[Tensor]]:
|
||||
config_split_ratio = CONFIG_PARSER.getfloat('training.loader', 'split_ratio')
|
||||
|
||||
dataset_size = len(dataset) # type:ignore[arg-type]
|
||||
training_size = int(dataset_size * config_split_ratio)
|
||||
validation_size = int(dataset_size - training_size)
|
||||
training_dataset, validate_dataset = random_split(dataset, [ training_size, validation_size ])
|
||||
return training_dataset, validate_dataset
|
||||
|
||||
|
||||
def create_trainer() -> Trainer:
|
||||
config_max_epochs = CONFIG_PARSER.getint('training.trainer', 'max_epochs')
|
||||
config_strategy = CONFIG_PARSER.get('training.trainer', 'strategy')
|
||||
config_precision = CONFIG_PARSER.get('training.trainer', 'precision')
|
||||
config_logger_path = CONFIG_PARSER.get('training.trainer', 'logger_path')
|
||||
config_logger_name = CONFIG_PARSER.get('training.trainer', 'logger_name')
|
||||
config_directory_path = CONFIG_PARSER.get('training.output', 'directory_path')
|
||||
config_file_pattern = CONFIG_PARSER.get('training.output', 'file_pattern')
|
||||
logger = TensorBoardLogger(config_logger_path, config_logger_name)
|
||||
|
||||
return Trainer(
|
||||
logger = logger,
|
||||
log_every_n_steps = 10,
|
||||
max_epochs = config_max_epochs,
|
||||
strategy = config_strategy,
|
||||
precision = config_precision,
|
||||
callbacks =
|
||||
[
|
||||
ModelCheckpoint(
|
||||
monitor = 'training_loss',
|
||||
dirpath = config_directory_path,
|
||||
filename = config_file_pattern,
|
||||
every_n_epochs = 1,
|
||||
save_top_k = 3,
|
||||
save_last = True
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def train() -> None:
|
||||
config_resume_path = CONFIG_PARSER.get('training.output', 'resume_path')
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.set_float32_matmul_precision('high')
|
||||
|
||||
dataset = StaticDataset(CONFIG_PARSER)
|
||||
training_loader, validation_loader = create_loaders(dataset)
|
||||
crossface_trainer = CrossFaceTrainer(CONFIG_PARSER)
|
||||
trainer = create_trainer()
|
||||
|
||||
if os.path.exists(config_resume_path):
|
||||
trainer.fit(crossface_trainer, training_loader, validation_loader, ckpt_path = config_resume_path)
|
||||
else:
|
||||
trainer.fit(crossface_trainer, training_loader, validation_loader)
|
||||
@@ -0,0 +1,8 @@
|
||||
from typing import Any, TypeAlias
|
||||
|
||||
from torch import Tensor
|
||||
|
||||
Batch : TypeAlias = Tensor
|
||||
Embedding : TypeAlias = Tensor
|
||||
|
||||
OptimizerSet : TypeAlias = Any
|
||||
@@ -0,0 +1,6 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
from src.training import train
|
||||
|
||||
if __name__ == '__main__':
|
||||
train()
|
||||
Reference in New Issue
Block a user