mirror of
https://github.com/facefusion/facefusion-labs.git
synced 2026-06-25 07:59:55 +02:00
changes
This commit is contained in:
@@ -30,12 +30,12 @@ class Discriminator(nn.Module):
|
||||
|
||||
return discriminators
|
||||
|
||||
def forward(self, input_tensor : Tensor) -> List[List[Tensor]]:
|
||||
def forward(self, input_tensor : Tensor) -> List[Tensor]:
|
||||
temp_tensor = input_tensor
|
||||
output_tensors = []
|
||||
|
||||
for discriminator in self.discriminators:
|
||||
output_tensors.append([ discriminator(temp_tensor) ])
|
||||
output_tensors.append(discriminator(temp_tensor))
|
||||
temp_tensor = self.avg_pool(temp_tensor)
|
||||
|
||||
return output_tensors
|
||||
|
||||
@@ -21,12 +21,12 @@ class DiscriminatorLoss(nn.Module):
|
||||
positive_tensors = []
|
||||
negative_tensors = []
|
||||
|
||||
for discriminator_output_tensor in discriminator_output_tensors:
|
||||
positive_tensor = torch.relu(discriminator_output_tensor[0] + 1).mean(dim = [ 1, 2, 3 ])
|
||||
for discriminator_source_tensor in discriminator_source_tensors:
|
||||
positive_tensor = torch.relu(discriminator_source_tensor + 1).mean(dim = [ 1, 2, 3 ])
|
||||
positive_tensors.append(positive_tensor)
|
||||
|
||||
for discriminator_source_tensor in discriminator_source_tensors:
|
||||
negative_tensor = torch.relu(1 - discriminator_source_tensor[0]).mean(dim = [ 1, 2, 3 ])
|
||||
for discriminator_output_tensor in discriminator_output_tensors:
|
||||
negative_tensor = torch.relu(1 - discriminator_output_tensor).mean(dim = [ 1, 2, 3 ])
|
||||
negative_tensors.append(negative_tensor)
|
||||
|
||||
positive_loss = torch.stack(positive_tensors).mean()
|
||||
@@ -44,7 +44,7 @@ class AdversarialLoss(nn.Module):
|
||||
temp_tensors = []
|
||||
|
||||
for discriminator_output_tensor in discriminator_output_tensors:
|
||||
temp_tensor = torch.relu(1 - discriminator_output_tensor[0]).mean(dim = [ 1, 2, 3 ]).mean()
|
||||
temp_tensor = torch.relu(1 - discriminator_output_tensor).mean(dim = [ 1, 2, 3 ]).mean()
|
||||
temp_tensors.append(temp_tensor)
|
||||
|
||||
adversarial_loss = torch.stack(temp_tensors).mean()
|
||||
@@ -164,8 +164,8 @@ class GazeLoss(nn.Module):
|
||||
return gaze_loss, weighted_gaze_loss
|
||||
|
||||
def detect_gaze(self, input_tensor : Tensor) -> Gaze:
|
||||
transform_size = CONFIG.getint('training.dataset', 'transform_size')
|
||||
crop_sizes = (torch.tensor([ 0.235, 0.875, 0.0625, 0.8 ]) * transform_size).int()
|
||||
output_size = CONFIG.getint('training.model.generator', 'output_size')
|
||||
crop_sizes = (torch.tensor([ 0.235, 0.875, 0.0625, 0.8 ]) * output_size).int()
|
||||
crop_tensor = input_tensor[:, :, crop_sizes[0]:crop_sizes[1], crop_sizes[2]:crop_sizes[3]]
|
||||
crop_tensor = (crop_tensor + 1) * 0.5
|
||||
crop_tensor = transforms.Normalize(mean = [ 0.485, 0.456, 0.406 ], std = [ 0.229, 0.224, 0.225 ])(crop_tensor)
|
||||
|
||||
@@ -198,15 +198,15 @@ def create_trainer() -> Trainer:
|
||||
def train() -> None:
|
||||
dataset_file_pattern = CONFIG.get('training.dataset', 'file_pattern')
|
||||
dataset_warp_template = cast(WarpTemplate, CONFIG.get('training.dataset', 'warp_template'))
|
||||
dataset_transform_size = CONFIG.getint('training.dataset', 'transform_size')
|
||||
dataset_batch_mode = cast(BatchMode, CONFIG.get('training.dataset', 'batch_mode'))
|
||||
dataset_batch_ratio = CONFIG.getfloat('training.dataset', 'batch_ratio')
|
||||
output_resume_path = CONFIG.get('training.output', 'resume_path')
|
||||
output_size = CONFIG.getint('training.model.generator', 'output_size')
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.set_float32_matmul_precision('high')
|
||||
|
||||
dataset = DynamicDataset(dataset_file_pattern, dataset_warp_template, dataset_transform_size, dataset_batch_mode, dataset_batch_ratio)
|
||||
dataset = DynamicDataset(dataset_file_pattern, dataset_warp_template, output_size, dataset_batch_mode, dataset_batch_ratio)
|
||||
training_loader, validation_loader = create_loaders(dataset)
|
||||
face_swapper_trainer = FaceSwapperTrainer()
|
||||
trainer = create_trainer()
|
||||
|
||||
Reference in New Issue
Block a user