mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-06-25 07:59:55 +02:00
More tweaks
This commit is contained in:
@@ -27,7 +27,7 @@ This `config.ini` utilizes the MegaFace dataset to train the Embedding Converter
|
||||
|
||||
```
|
||||
[training.dataset]
|
||||
file_pattern = .datasets/images/{}/*.*g
|
||||
file_pattern = .datasets/megaface/**/*.jpg
|
||||
```
|
||||
|
||||
```
|
||||
|
||||
@@ -11,17 +11,17 @@ from .types import Batch
|
||||
|
||||
|
||||
class DynamicDataset(Dataset[Tensor]):
|
||||
def __init__(self, dataset_file_pattern : str) -> None:
|
||||
self.image_paths = glob.glob(dataset_file_pattern)
|
||||
def __init__(self, file_pattern : str) -> None:
|
||||
self.file_paths = glob.glob(file_pattern)
|
||||
self.transforms = self.compose_transforms()
|
||||
|
||||
def __getitem__(self, index : int) -> Batch:
|
||||
image_path = random.choice(self.image_paths)
|
||||
image_path = random.choice(self.file_paths)
|
||||
vision_frame = cv2.imread(image_path)
|
||||
return self.transforms(vision_frame)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.image_paths)
|
||||
return len(self.file_paths)
|
||||
|
||||
@staticmethod
|
||||
def compose_transforms() -> transforms:
|
||||
|
||||
@@ -24,13 +24,15 @@ class EmbeddingConverterTrainer(lightning.LightningModule):
|
||||
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')
|
||||
learning_rate = 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()
|
||||
self.lr = learning_rate
|
||||
|
||||
def forward(self, source_embedding : Embedding) -> Embedding:
|
||||
return self.embedding_converter(source_embedding)
|
||||
@@ -84,8 +86,8 @@ def create_loaders(dataset : Dataset[Tensor]) -> Tuple[DataLoader[Tensor], DataL
|
||||
def split_dataset(dataset : Dataset[Tensor]) -> Tuple[Dataset[Tensor], Dataset[Tensor]]:
|
||||
loader_split_ratio = CONFIG.getfloat('training.loader', 'split_ratio')
|
||||
dataset_size = len(dataset) # type:ignore[arg-type]
|
||||
training_size = dataset_size * loader_split_ratio
|
||||
validation_size = dataset_size - training_size
|
||||
training_size = int(dataset_size * loader_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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user