This commit is contained in:
harisreedhar
2025-03-06 15:00:42 +05:30
committed by henryruhs
parent dfd018a897
commit 61f48d9246
3 changed files with 11 additions and 11 deletions
+2 -2
View File
@@ -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
+7 -7
View File
@@ -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)
+2 -2
View File
@@ -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()